Add Supabase tenant extension as built-in (#267)
Move the Supabase tenant extension into the hindsight-api package so users can enable it with just an environment variable — no file copying or Docker image modifications needed. Key improvements over the original submission: - JWKS-based local JWT verification (no network call per request) with automatic fallback to /auth/v1/user for legacy HS256 projects - Service key is now optional (only needed for HS256 or health checks) - UUID validation on user IDs before schema name construction - Schema prefix validation against Postgres identifier rules - Key rotation handling with automatic JWKS cache refresh - Proper logging via Python logging module - Tenant extension lifecycle hooks (on_startup/on_shutdown) wired into the server lifespan - Public tenant_extension property on MemoryEngine - 54 unit tests covering both verification modes, cache behavior, error paths, and the extension loader - README updated to reflect JWKS-first architecture Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
parent
c568094b8c
commit
e99ee0f243
9 changed files with 4295 additions and 3196 deletions
|
|
@ -1432,6 +1432,12 @@ def create_app(
|
|||
poller_task = asyncio.create_task(poller.run())
|
||||
logging.info(f"Worker poller started (worker_id={worker_id})")
|
||||
|
||||
# Call tenant extension startup hook (e.g. JWKS fetch for Supabase)
|
||||
tenant_extension = memory.tenant_extension
|
||||
if tenant_extension:
|
||||
await tenant_extension.on_startup()
|
||||
logging.info("Tenant extension started")
|
||||
|
||||
# Call HTTP extension startup hook
|
||||
if http_extension:
|
||||
await http_extension.on_startup()
|
||||
|
|
@ -1450,6 +1456,11 @@ def create_app(
|
|||
pass
|
||||
logging.info("Worker poller stopped")
|
||||
|
||||
# Call tenant extension shutdown hook
|
||||
if tenant_extension:
|
||||
await tenant_extension.on_shutdown()
|
||||
logging.info("Tenant extension stopped")
|
||||
|
||||
# Call HTTP extension shutdown hook
|
||||
if http_extension:
|
||||
await http_extension.on_shutdown()
|
||||
|
|
|
|||
|
|
@ -467,6 +467,11 @@ class MemoryEngine(MemoryEngineInterface):
|
|||
tenant_extension = DefaultTenantExtension(config={})
|
||||
self._tenant_extension = tenant_extension
|
||||
|
||||
@property
|
||||
def tenant_extension(self) -> "TenantExtension | None":
|
||||
"""The configured tenant extension, if any."""
|
||||
return self._tenant_extension
|
||||
|
||||
async def _validate_operation(self, validation_coro) -> None:
|
||||
"""
|
||||
Run validation if an operation validator is configured.
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ with the system (e.g., running migrations for tenant schemas).
|
|||
"""
|
||||
|
||||
from hindsight_api.extensions.base import Extension
|
||||
from hindsight_api.extensions.builtin import ApiKeyTenantExtension
|
||||
from hindsight_api.extensions.builtin import ApiKeyTenantExtension, SupabaseTenantExtension
|
||||
from hindsight_api.extensions.context import DefaultExtensionContext, ExtensionContext
|
||||
from hindsight_api.extensions.http import HttpExtension
|
||||
from hindsight_api.extensions.loader import load_extension
|
||||
|
|
@ -80,6 +80,7 @@ __all__ = [
|
|||
"MentalModelRefreshResult",
|
||||
# Tenant/Auth
|
||||
"ApiKeyTenantExtension",
|
||||
"SupabaseTenantExtension",
|
||||
"AuthenticationError",
|
||||
"RequestContext",
|
||||
"Tenant",
|
||||
|
|
|
|||
|
|
@ -6,13 +6,17 @@ They can be used directly or serve as examples for custom implementations.
|
|||
|
||||
Available built-in extensions:
|
||||
- ApiKeyTenantExtension: Simple API key validation with public schema
|
||||
- SupabaseTenantExtension: Supabase JWT validation with per-user schema isolation
|
||||
|
||||
Example usage:
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.tenant:ApiKeyTenantExtension
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension
|
||||
"""
|
||||
|
||||
from hindsight_api.extensions.builtin.supabase_tenant import SupabaseTenantExtension
|
||||
from hindsight_api.extensions.builtin.tenant import ApiKeyTenantExtension
|
||||
|
||||
__all__ = [
|
||||
"ApiKeyTenantExtension",
|
||||
"SupabaseTenantExtension",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,423 @@
|
|||
"""
|
||||
Supabase Tenant Extension for Hindsight
|
||||
|
||||
Validates Supabase JWTs and maps authenticated users to isolated memory banks.
|
||||
Each user gets their own PostgreSQL schema based on their Supabase user ID.
|
||||
|
||||
This extension enables multi-tenant memory isolation for applications using
|
||||
Supabase Auth - each authenticated user's memories are stored in a separate
|
||||
schema, ensuring complete data isolation.
|
||||
|
||||
JWT Verification Strategy:
|
||||
By default, JWTs are verified locally using public keys from the Supabase
|
||||
JWKS endpoint (/auth/v1/.well-known/jwks.json). This is the Supabase-recommended
|
||||
approach: no network call per request, fast, and secure.
|
||||
|
||||
If JWKS keys are unavailable (e.g., legacy HS256 projects), the extension
|
||||
falls back to calling /auth/v1/user per request for validation. This requires
|
||||
the service_role key to be configured.
|
||||
|
||||
Configuration via environment variables:
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension
|
||||
HINDSIGHT_API_TENANT_SUPABASE_URL=https://your-project.supabase.co
|
||||
|
||||
# Optional - only required for legacy HS256 projects or health checks
|
||||
HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY=your-service-role-key
|
||||
|
||||
# Optional
|
||||
HINDSIGHT_API_TENANT_SCHEMA_PREFIX=user # Default: "user" (creates user_<uuid> schemas)
|
||||
|
||||
Usage:
|
||||
Clients pass their Supabase JWT in the Authorization header:
|
||||
|
||||
curl -H "Authorization: Bearer <supabase_jwt>" \\
|
||||
https://your-hindsight-server/v1/default/banks/my-bank/memories/recall
|
||||
|
||||
Author: BrighterBalance (https://brighterbalance.app)
|
||||
License: MIT
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
|
||||
import httpx
|
||||
import jwt as pyjwt
|
||||
from jwt import PyJWK
|
||||
|
||||
from hindsight_api.extensions.tenant import AuthenticationError, Tenant, TenantContext, TenantExtension
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
__all__ = ["SupabaseTenantExtension"]
|
||||
|
||||
# Minimum expected JWT length (JWTs are typically 100+ characters)
|
||||
MIN_TOKEN_LENGTH = 20
|
||||
|
||||
# Timeout for Supabase API calls
|
||||
REQUEST_TIMEOUT_SECONDS = 10.0
|
||||
|
||||
# JWKS cache TTL — Supabase Edge caches JWKS for 10 minutes, so we match that
|
||||
JWKS_CACHE_TTL_SECONDS = 600
|
||||
|
||||
# Minimum interval between JWKS refreshes to avoid hammering the endpoint
|
||||
JWKS_MIN_REFRESH_INTERVAL_SECONDS = 30
|
||||
|
||||
# Algorithms supported by Supabase Auth for asymmetric JWT signing
|
||||
SUPPORTED_ALGORITHMS = ["RS256", "ES256"]
|
||||
|
||||
# Supabase user IDs are UUIDs — validate before using in schema names
|
||||
_UUID_RE = re.compile(r"^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$", re.IGNORECASE)
|
||||
|
||||
# Schema prefix must be a valid Postgres identifier component (letters, digits, underscores)
|
||||
_SCHEMA_PREFIX_RE = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$")
|
||||
|
||||
|
||||
class SupabaseTenantExtension(TenantExtension):
|
||||
"""
|
||||
TenantExtension that validates Supabase JWTs for multi-tenant isolation.
|
||||
|
||||
Each authenticated user gets their own PostgreSQL schema, ensuring complete
|
||||
memory isolation between users. The schema name is derived from the user's
|
||||
Supabase user ID (the ``sub`` claim in the JWT).
|
||||
|
||||
JWT verification uses JWKS (local, no network call per request) when
|
||||
asymmetric keys are configured in Supabase, and falls back to the
|
||||
``/auth/v1/user`` endpoint for legacy HS256 projects.
|
||||
|
||||
Example:
|
||||
User with ID "a1b2c3d4-e5f6-7890-abcd-ef1234567890"
|
||||
gets schema "user_a1b2c3d4_e5f6_7890_abcd_ef1234567890"
|
||||
"""
|
||||
|
||||
def __init__(self, config: dict[str, str]) -> None:
|
||||
"""
|
||||
Initialize with configuration from environment variables.
|
||||
|
||||
Config keys are derived from HINDSIGHT_API_TENANT_* env vars:
|
||||
- HINDSIGHT_API_TENANT_SUPABASE_URL -> config["supabase_url"] (required)
|
||||
- HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY -> config["supabase_service_key"] (optional)
|
||||
- HINDSIGHT_API_TENANT_SCHEMA_PREFIX -> config["schema_prefix"] (optional)
|
||||
|
||||
Args:
|
||||
config: Dictionary of configuration values from environment
|
||||
|
||||
Raises:
|
||||
ValueError: If required configuration is missing
|
||||
"""
|
||||
super().__init__(config)
|
||||
|
||||
self.supabase_url = (config.get("supabase_url") or "").rstrip("/")
|
||||
self.supabase_service_key = config.get("supabase_service_key")
|
||||
self.schema_prefix = config.get("schema_prefix", "user")
|
||||
|
||||
# Track initialized schemas to avoid redundant migrations
|
||||
self._initialized_schemas: set[str] = set()
|
||||
|
||||
# Reusable HTTP client (created on startup)
|
||||
self._http_client: httpx.AsyncClient | None = None
|
||||
|
||||
# JWKS state
|
||||
self._jwks_keys: dict[str, PyJWK] = {}
|
||||
self._jwks_last_fetched: float = 0
|
||||
self._use_jwks: bool = False
|
||||
|
||||
if not self.supabase_url:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_TENANT_SUPABASE_URL is required. "
|
||||
"Set it to your Supabase project URL (e.g., https://xxx.supabase.co)"
|
||||
)
|
||||
|
||||
if not _SCHEMA_PREFIX_RE.match(self.schema_prefix):
|
||||
raise ValueError(
|
||||
f"Invalid schema_prefix '{self.schema_prefix}'. "
|
||||
"Must be a valid Postgres identifier (letters, digits, underscores, starting with a letter or underscore)."
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Lifecycle
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def on_startup(self) -> None:
|
||||
"""
|
||||
Called when Hindsight starts.
|
||||
|
||||
Creates a reusable HTTP client, fetches JWKS for local JWT verification,
|
||||
and optionally verifies connectivity to Supabase.
|
||||
"""
|
||||
logger.info("Initializing Supabase tenant extension")
|
||||
logger.info("Supabase URL: %s", self.supabase_url)
|
||||
logger.info("Schema prefix: %s_", self.schema_prefix)
|
||||
|
||||
self._http_client = httpx.AsyncClient(timeout=REQUEST_TIMEOUT_SECONDS)
|
||||
|
||||
# Attempt to fetch JWKS for fast local JWT verification
|
||||
await self._try_init_jwks()
|
||||
|
||||
# Optional health check using service key
|
||||
if self.supabase_service_key:
|
||||
await self._health_check()
|
||||
|
||||
async def on_shutdown(self) -> None:
|
||||
"""Called when Hindsight shuts down. Closes the HTTP client."""
|
||||
logger.info("Shutting down Supabase tenant extension")
|
||||
if self._http_client:
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# JWKS management
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _try_init_jwks(self) -> None:
|
||||
"""Fetch JWKS and decide verification mode (local JWKS vs legacy endpoint)."""
|
||||
try:
|
||||
await self._fetch_jwks()
|
||||
if self._jwks_keys:
|
||||
self._use_jwks = True
|
||||
logger.info(
|
||||
"JWKS loaded — using local JWT verification with %d key(s)",
|
||||
len(self._jwks_keys),
|
||||
)
|
||||
return
|
||||
|
||||
# JWKS endpoint returned no keys — project likely uses legacy HS256
|
||||
logger.warning(
|
||||
"JWKS endpoint returned no signing keys. "
|
||||
"Falling back to /auth/v1/user endpoint for JWT verification. "
|
||||
"For better performance, enable asymmetric JWT signing in your "
|
||||
"Supabase dashboard (Project Settings → Auth → JWT Algorithm)."
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Could not fetch JWKS (%s). Falling back to /auth/v1/user endpoint for JWT verification.",
|
||||
e,
|
||||
)
|
||||
|
||||
# Legacy mode requires service key
|
||||
if not self.supabase_service_key:
|
||||
raise ValueError(
|
||||
"HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY is required when JWKS "
|
||||
"is not available. Either enable asymmetric JWT signing in your "
|
||||
"Supabase project or provide the service_role key."
|
||||
)
|
||||
self._use_jwks = False
|
||||
|
||||
async def _fetch_jwks(self) -> None:
|
||||
"""Fetch public signing keys from the Supabase JWKS endpoint."""
|
||||
if self._http_client is None:
|
||||
raise RuntimeError("HTTP client not initialized")
|
||||
|
||||
url = f"{self.supabase_url}/auth/v1/.well-known/jwks.json"
|
||||
response = await self._http_client.get(url)
|
||||
response.raise_for_status()
|
||||
|
||||
jwks_data = response.json()
|
||||
keys: dict[str, PyJWK] = {}
|
||||
for key_data in jwks_data.get("keys", []):
|
||||
kid = key_data.get("kid")
|
||||
if kid:
|
||||
keys[kid] = PyJWK(key_data)
|
||||
|
||||
self._jwks_keys = keys
|
||||
self._jwks_last_fetched = time.monotonic()
|
||||
|
||||
async def _get_signing_key(self, token: str) -> PyJWK:
|
||||
"""
|
||||
Resolve the signing key for a token from the JWKS cache.
|
||||
|
||||
If the key ID (``kid``) is not in the cache, triggers one JWKS refresh
|
||||
to handle key rotation before raising an error.
|
||||
"""
|
||||
header = pyjwt.get_unverified_header(token)
|
||||
kid = header.get("kid")
|
||||
if not kid:
|
||||
raise AuthenticationError("Token missing key ID (kid) header")
|
||||
|
||||
# Refresh cache if stale
|
||||
now = time.monotonic()
|
||||
if now - self._jwks_last_fetched > JWKS_CACHE_TTL_SECONDS:
|
||||
logger.debug("JWKS cache expired, refreshing")
|
||||
await self._fetch_jwks()
|
||||
|
||||
if kid in self._jwks_keys:
|
||||
return self._jwks_keys[kid]
|
||||
|
||||
# Key not found — try one forced refresh to handle key rotation,
|
||||
# but only if we haven't just refreshed
|
||||
if now - self._jwks_last_fetched > JWKS_MIN_REFRESH_INTERVAL_SECONDS:
|
||||
logger.info("Signing key %s not in cache, refreshing JWKS for possible key rotation", kid)
|
||||
await self._fetch_jwks()
|
||||
if kid in self._jwks_keys:
|
||||
return self._jwks_keys[kid]
|
||||
|
||||
raise AuthenticationError("Unable to find signing key for token")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Authentication
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def authenticate(self, context: RequestContext) -> TenantContext:
|
||||
"""
|
||||
Validate a Supabase JWT and return tenant context.
|
||||
|
||||
Uses local JWKS verification when available (no network call per
|
||||
request), falling back to the ``/auth/v1/user`` endpoint for legacy
|
||||
HS256 projects.
|
||||
|
||||
Args:
|
||||
context: Request context containing the API key (JWT)
|
||||
|
||||
Returns:
|
||||
TenantContext with schema_name set to ``{prefix}_{user_uuid}``
|
||||
|
||||
Raises:
|
||||
AuthenticationError: If token is missing, invalid, or expired
|
||||
"""
|
||||
token = context.api_key
|
||||
|
||||
if not token:
|
||||
raise AuthenticationError("Missing Authorization header. Expected: Bearer <supabase_jwt>")
|
||||
|
||||
if len(token) < MIN_TOKEN_LENGTH:
|
||||
raise AuthenticationError("Invalid token format")
|
||||
|
||||
if self._http_client is None:
|
||||
raise AuthenticationError("Extension not initialized")
|
||||
|
||||
# Verify the JWT and extract user ID
|
||||
if self._use_jwks:
|
||||
user_id = await self._verify_token_jwks(token)
|
||||
else:
|
||||
user_id = await self._verify_token_legacy(token)
|
||||
|
||||
# Validate user ID format before using in schema name
|
||||
if not _UUID_RE.match(user_id):
|
||||
raise AuthenticationError("Invalid user ID format in token")
|
||||
|
||||
# Build isolated schema name — hyphens to underscores for Postgres compatibility
|
||||
safe_user_id = user_id.replace("-", "_")
|
||||
schema_name = f"{self.schema_prefix}_{safe_user_id}"
|
||||
|
||||
# Initialize schema on first access
|
||||
if schema_name not in self._initialized_schemas:
|
||||
await self._initialize_schema(schema_name)
|
||||
|
||||
return TenantContext(schema_name=schema_name)
|
||||
|
||||
async def _verify_token_jwks(self, token: str) -> str:
|
||||
"""
|
||||
Verify a JWT locally using cached JWKS public keys.
|
||||
|
||||
Validates signature, expiration, issuer, and audience. Returns the
|
||||
user ID from the ``sub`` claim.
|
||||
|
||||
Raises:
|
||||
AuthenticationError: If the token is invalid or expired.
|
||||
"""
|
||||
try:
|
||||
signing_key = await self._get_signing_key(token)
|
||||
payload = pyjwt.decode(
|
||||
token,
|
||||
signing_key.key,
|
||||
algorithms=SUPPORTED_ALGORITHMS,
|
||||
audience="authenticated",
|
||||
issuer=f"{self.supabase_url}/auth/v1",
|
||||
)
|
||||
except pyjwt.ExpiredSignatureError:
|
||||
raise AuthenticationError("Token has expired")
|
||||
except pyjwt.InvalidAudienceError:
|
||||
raise AuthenticationError("Invalid token audience")
|
||||
except pyjwt.InvalidIssuerError:
|
||||
raise AuthenticationError("Invalid token issuer")
|
||||
except pyjwt.DecodeError:
|
||||
raise AuthenticationError("Invalid token")
|
||||
except AuthenticationError:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise AuthenticationError(f"Token verification failed: {e!s}")
|
||||
|
||||
user_id = payload.get("sub")
|
||||
if not user_id:
|
||||
raise AuthenticationError("Token valid but missing subject (sub) claim")
|
||||
return user_id
|
||||
|
||||
async def _verify_token_legacy(self, token: str) -> str:
|
||||
"""
|
||||
Verify a JWT by calling the Supabase ``/auth/v1/user`` endpoint.
|
||||
|
||||
This is the fallback for projects using legacy HS256 JWT signing.
|
||||
Adds a network round-trip per request.
|
||||
|
||||
Raises:
|
||||
AuthenticationError: If the token is invalid or the request fails.
|
||||
"""
|
||||
try:
|
||||
response = await self._http_client.get(
|
||||
f"{self.supabase_url}/auth/v1/user",
|
||||
headers={
|
||||
"Authorization": f"Bearer {token}",
|
||||
"apikey": self.supabase_service_key,
|
||||
},
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise AuthenticationError("Invalid or expired token")
|
||||
|
||||
if response.status_code != 200:
|
||||
raise AuthenticationError(f"Authentication failed: {response.status_code}")
|
||||
|
||||
user_data = response.json()
|
||||
user_id = user_data.get("id")
|
||||
|
||||
if not user_id:
|
||||
raise AuthenticationError("Token valid but no user ID found")
|
||||
|
||||
return user_id
|
||||
|
||||
except AuthenticationError:
|
||||
raise
|
||||
except httpx.TimeoutException:
|
||||
raise AuthenticationError("Authentication timeout - please retry")
|
||||
except httpx.RequestError as e:
|
||||
raise AuthenticationError(f"Connection error: {e!s}")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Schema management
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _initialize_schema(self, schema_name: str) -> None:
|
||||
"""Run migrations for a new tenant schema and cache the result."""
|
||||
logger.info("Initializing schema: %s", schema_name)
|
||||
try:
|
||||
await self.context.run_migration(schema_name)
|
||||
self._initialized_schemas.add(schema_name)
|
||||
logger.info("Schema ready: %s", schema_name)
|
||||
except Exception as e:
|
||||
logger.error("Schema initialization failed for %s: %s", schema_name, e)
|
||||
raise AuthenticationError(f"Failed to initialize tenant: {e!s}")
|
||||
|
||||
async def list_tenants(self) -> list[Tenant]:
|
||||
"""Return all tenant schemas that have been initialized."""
|
||||
return [Tenant(schema=schema) for schema in self._initialized_schemas]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Health check
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _health_check(self) -> None:
|
||||
"""Verify connectivity to Supabase using the auth health endpoint."""
|
||||
try:
|
||||
response = await self._http_client.get(
|
||||
f"{self.supabase_url}/auth/v1/health",
|
||||
headers={"apikey": self.supabase_service_key},
|
||||
)
|
||||
if response.status_code == 200:
|
||||
logger.info("Supabase connection verified")
|
||||
else:
|
||||
logger.warning("Supabase health check returned %d", response.status_code)
|
||||
except Exception as e:
|
||||
logger.warning("Could not verify Supabase connection: %s", e)
|
||||
|
|
@ -25,6 +25,7 @@ dependencies = [
|
|||
"psycopg2-binary>=2.9.11",
|
||||
"tiktoken>=0.12.0",
|
||||
"httpx>=0.27.0",
|
||||
"PyJWT[crypto]>=2.8.0",
|
||||
"fastmcp>=2.14.0", # CVE-2025-66416
|
||||
"pg0-embedded>=0.11.0",
|
||||
"python-dateutil>=2.8.0",
|
||||
|
|
|
|||
834
hindsight-api/tests/test_supabase_tenant.py
Normal file
834
hindsight-api/tests/test_supabase_tenant.py
Normal file
|
|
@ -0,0 +1,834 @@
|
|||
"""Tests for the Supabase Tenant Extension."""
|
||||
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import jwt as pyjwt
|
||||
import pytest
|
||||
from jwt import PyJWK
|
||||
|
||||
from hindsight_api.extensions.builtin.supabase_tenant import (
|
||||
JWKS_CACHE_TTL_SECONDS,
|
||||
JWKS_MIN_REFRESH_INTERVAL_SECONDS,
|
||||
MIN_TOKEN_LENGTH,
|
||||
SupabaseTenantExtension,
|
||||
)
|
||||
from hindsight_api.extensions.context import ExtensionContext
|
||||
from hindsight_api.extensions.loader import load_extension
|
||||
from hindsight_api.extensions.tenant import AuthenticationError, Tenant, TenantContext, TenantExtension
|
||||
from hindsight_api.models import RequestContext
|
||||
|
||||
# A valid UUID for test user IDs
|
||||
VALID_UUID = "a1b2c3d4-e5f6-7890-abcd-ef1234567890"
|
||||
|
||||
# Minimal JWKS response with one RSA key
|
||||
MOCK_JWKS_RESPONSE = {
|
||||
"keys": [
|
||||
{
|
||||
"kid": "test-key-1",
|
||||
"kty": "RSA",
|
||||
"alg": "RS256",
|
||||
"use": "sig",
|
||||
"n": "0vx7agoebGcQSuuPiLJXZptN9nndrQmbXEps2aiAFbWhM78LhWx4cbbfAAtVT86zwu1RK7aPFFxuhDR1L6tSoc_BJECPebWKRXjBZCiFV4n3oknjhMstn64tZ_2W-5JsGY4Hc5n9yBXArwl93lqt7_RN5w6Cf0h4QyQ5v-65YGjQR0_FDW2QvzqY368QQMicAtaSqzs8KJZgnYb9c7d0zgdAZHzu6qMQvRL5hajrn1n91CbOpbISD08qNLyrdkt-bFTWhAI4vMQFh6WeZu0fM4lFd2NcRwr3XPksINHaQ-G_xBniIqbw0Ls1jF44-csFCur-kEgU8awapJzKnqDKgw",
|
||||
"e": "AQAB",
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _make_extension(
|
||||
supabase_url: str = "https://test.supabase.co",
|
||||
service_key: str | None = "test-service-key",
|
||||
schema_prefix: str | None = None,
|
||||
) -> SupabaseTenantExtension:
|
||||
"""Helper to create a SupabaseTenantExtension with test config."""
|
||||
config = {
|
||||
"supabase_url": supabase_url,
|
||||
}
|
||||
if service_key is not None:
|
||||
config["supabase_service_key"] = service_key
|
||||
if schema_prefix is not None:
|
||||
config["schema_prefix"] = schema_prefix
|
||||
return SupabaseTenantExtension(config)
|
||||
|
||||
|
||||
def _make_mock_response(status_code: int = 200, json_data: dict | None = None) -> MagicMock:
|
||||
"""Helper to create a mock httpx.Response."""
|
||||
response = MagicMock(spec=httpx.Response)
|
||||
response.status_code = status_code
|
||||
response.json.return_value = json_data or {}
|
||||
response.raise_for_status = MagicMock()
|
||||
if status_code >= 400:
|
||||
response.raise_for_status.side_effect = httpx.HTTPStatusError("error", request=MagicMock(), response=response)
|
||||
return response
|
||||
|
||||
|
||||
def _make_valid_token() -> str:
|
||||
"""Return a token that passes the MIN_TOKEN_LENGTH check."""
|
||||
return "a" * (MIN_TOKEN_LENGTH + 10)
|
||||
|
||||
|
||||
def _setup_jwks_ext() -> tuple[SupabaseTenantExtension, AsyncMock]:
|
||||
"""Create an extension in JWKS mode with mocked internals."""
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
ext._http_client = mock_client
|
||||
ext._use_jwks = True
|
||||
ext._jwks_keys = {"test-key-1": MagicMock(spec=PyJWK)}
|
||||
ext._jwks_keys["test-key-1"].key = "mock-public-key"
|
||||
ext._jwks_last_fetched = time.monotonic()
|
||||
return ext, mock_client
|
||||
|
||||
|
||||
def _setup_legacy_ext() -> tuple[SupabaseTenantExtension, AsyncMock]:
|
||||
"""Create an extension in legacy mode with mocked internals."""
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
ext._http_client = mock_client
|
||||
ext._use_jwks = False
|
||||
return ext, mock_client
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Initialization
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestSupabaseTenantExtensionInit:
|
||||
"""Tests for extension initialization."""
|
||||
|
||||
def test_init_with_valid_config(self):
|
||||
ext = _make_extension()
|
||||
assert ext.supabase_url == "https://test.supabase.co"
|
||||
assert ext.supabase_service_key == "test-service-key"
|
||||
assert ext.schema_prefix == "user"
|
||||
assert ext._initialized_schemas == set()
|
||||
assert ext._http_client is None
|
||||
assert ext._use_jwks is False
|
||||
assert ext._jwks_keys == {}
|
||||
|
||||
def test_init_missing_supabase_url(self):
|
||||
with pytest.raises(ValueError, match="HINDSIGHT_API_TENANT_SUPABASE_URL is required"):
|
||||
SupabaseTenantExtension({})
|
||||
|
||||
def test_init_without_service_key(self):
|
||||
"""Service key is optional — JWKS mode doesn't require it."""
|
||||
ext = _make_extension(service_key=None)
|
||||
assert ext.supabase_service_key is None
|
||||
|
||||
def test_init_default_schema_prefix(self):
|
||||
ext = _make_extension()
|
||||
assert ext.schema_prefix == "user"
|
||||
|
||||
def test_init_custom_schema_prefix(self):
|
||||
ext = _make_extension(schema_prefix="tenant")
|
||||
assert ext.schema_prefix == "tenant"
|
||||
|
||||
def test_init_strips_trailing_slash(self):
|
||||
ext = _make_extension(supabase_url="https://test.supabase.co/")
|
||||
assert ext.supabase_url == "https://test.supabase.co"
|
||||
|
||||
def test_init_rejects_invalid_schema_prefix(self):
|
||||
"""Schema prefix with special characters should be rejected."""
|
||||
with pytest.raises(ValueError, match="Invalid schema_prefix"):
|
||||
_make_extension(schema_prefix='"; DROP TABLE')
|
||||
|
||||
def test_init_rejects_empty_schema_prefix(self):
|
||||
with pytest.raises(ValueError, match="Invalid schema_prefix"):
|
||||
_make_extension(schema_prefix="")
|
||||
|
||||
def test_init_rejects_schema_prefix_starting_with_digit(self):
|
||||
with pytest.raises(ValueError, match="Invalid schema_prefix"):
|
||||
_make_extension(schema_prefix="123abc")
|
||||
|
||||
def test_init_allows_underscore_prefix(self):
|
||||
ext = _make_extension(schema_prefix="_internal")
|
||||
assert ext.schema_prefix == "_internal"
|
||||
|
||||
def test_is_tenant_extension_subclass(self):
|
||||
ext = _make_extension()
|
||||
assert isinstance(ext, TenantExtension)
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Startup — JWKS initialization
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestSupabaseTenantExtensionStartup:
|
||||
"""Tests for on_startup behavior."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_startup_creates_http_client(self):
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
# JWKS fetch returns keys
|
||||
mock_client.get.return_value = _make_mock_response(200, MOCK_JWKS_RESPONSE)
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.httpx.AsyncClient", return_value=mock_client):
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.PyJWK"):
|
||||
await ext.on_startup()
|
||||
|
||||
assert ext._http_client is mock_client
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_startup_fetches_jwks(self):
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
mock_client.get.return_value = _make_mock_response(200, MOCK_JWKS_RESPONSE)
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.httpx.AsyncClient", return_value=mock_client):
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.PyJWK") as mock_pyjwk:
|
||||
mock_pyjwk.return_value = MagicMock(spec=PyJWK)
|
||||
await ext.on_startup()
|
||||
|
||||
assert ext._use_jwks is True
|
||||
# First call: JWKS fetch, second call: health check
|
||||
assert mock_client.get.call_count == 2
|
||||
jwks_call = mock_client.get.call_args_list[0]
|
||||
assert jwks_call.args[0] == "https://test.supabase.co/auth/v1/.well-known/jwks.json"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_startup_falls_back_to_legacy_when_jwks_empty(self):
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
|
||||
# JWKS returns empty keys, health check succeeds
|
||||
def mock_get(url, **kwargs):
|
||||
if "jwks" in url:
|
||||
return _make_mock_response(200, {"keys": []})
|
||||
return _make_mock_response(200)
|
||||
|
||||
mock_client.get.side_effect = mock_get
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.httpx.AsyncClient", return_value=mock_client):
|
||||
await ext.on_startup()
|
||||
|
||||
assert ext._use_jwks is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_startup_falls_back_to_legacy_when_jwks_fetch_fails(self):
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
|
||||
call_count = 0
|
||||
|
||||
def mock_get(url, **kwargs):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
# JWKS fetch fails
|
||||
raise httpx.ConnectError("Connection refused")
|
||||
# health check
|
||||
return _make_mock_response(200)
|
||||
|
||||
mock_client.get.side_effect = mock_get
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.httpx.AsyncClient", return_value=mock_client):
|
||||
await ext.on_startup()
|
||||
|
||||
assert ext._use_jwks is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_startup_raises_if_no_jwks_and_no_service_key(self):
|
||||
ext = _make_extension(service_key=None)
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
mock_client.get.return_value = _make_mock_response(200, {"keys": []})
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.httpx.AsyncClient", return_value=mock_client):
|
||||
with pytest.raises(ValueError, match="HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY is required"):
|
||||
await ext.on_startup()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_startup_health_check_with_service_key(self):
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
mock_client.get.return_value = _make_mock_response(200, MOCK_JWKS_RESPONSE)
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.httpx.AsyncClient", return_value=mock_client):
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.PyJWK"):
|
||||
await ext.on_startup()
|
||||
|
||||
# Second call should be health check
|
||||
health_call = mock_client.get.call_args_list[1]
|
||||
assert health_call.args[0] == "https://test.supabase.co/auth/v1/health"
|
||||
assert health_call.kwargs["headers"] == {"apikey": "test-service-key"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_startup_skips_health_check_without_service_key(self):
|
||||
ext = _make_extension(service_key=None)
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
mock_client.get.return_value = _make_mock_response(200, MOCK_JWKS_RESPONSE)
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.httpx.AsyncClient", return_value=mock_client):
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.PyJWK"):
|
||||
await ext.on_startup()
|
||||
|
||||
# Only one call: JWKS fetch, no health check
|
||||
assert mock_client.get.call_count == 1
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# JWKS cache management
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestJWKSCacheManagement:
|
||||
"""Tests for JWKS key fetching, caching, and rotation handling."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_signing_key_from_cache(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header:
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
key = await ext._get_signing_key("fake-token")
|
||||
|
||||
assert key is ext._jwks_keys["test-key-1"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_signing_key_refreshes_stale_cache(self):
|
||||
ext, mock_client = _setup_jwks_ext()
|
||||
# Make cache expired
|
||||
ext._jwks_last_fetched = time.monotonic() - JWKS_CACHE_TTL_SECONDS - 1
|
||||
|
||||
new_key = MagicMock(spec=PyJWK)
|
||||
mock_client.get.return_value = _make_mock_response(200, MOCK_JWKS_RESPONSE)
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.PyJWK", return_value=new_key),
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
key = await ext._get_signing_key("fake-token")
|
||||
|
||||
assert key is new_key
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_signing_key_handles_key_rotation(self):
|
||||
"""When kid not in cache and cache is old enough, refresh once for key rotation."""
|
||||
ext, mock_client = _setup_jwks_ext()
|
||||
# Make cache just old enough to allow a refresh
|
||||
ext._jwks_last_fetched = time.monotonic() - JWKS_MIN_REFRESH_INTERVAL_SECONDS - 1
|
||||
|
||||
rotated_key = MagicMock(spec=PyJWK)
|
||||
mock_client.get.return_value = _make_mock_response(200, MOCK_JWKS_RESPONSE)
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.PyJWK", return_value=rotated_key),
|
||||
):
|
||||
mock_header.return_value = {"kid": "rotated-key-99", "alg": "RS256"}
|
||||
# The refreshed JWKS won't have "rotated-key-99" either, so this should raise
|
||||
with pytest.raises(AuthenticationError, match="Unable to find signing key"):
|
||||
await ext._get_signing_key("fake-token")
|
||||
|
||||
# Should have attempted one refresh
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_signing_key_missing_kid_header(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header:
|
||||
mock_header.return_value = {"alg": "RS256"} # no kid
|
||||
with pytest.raises(AuthenticationError, match="Token missing key ID"):
|
||||
await ext._get_signing_key("fake-token")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_signing_key_refresh_network_error(self):
|
||||
"""If JWKS refresh fails during key rotation, error should propagate."""
|
||||
ext, mock_client = _setup_jwks_ext()
|
||||
ext._jwks_last_fetched = time.monotonic() - JWKS_MIN_REFRESH_INTERVAL_SECONDS - 1
|
||||
|
||||
mock_client.get.side_effect = httpx.ConnectError("Connection refused")
|
||||
|
||||
with patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header:
|
||||
mock_header.return_value = {"kid": "unknown-key", "alg": "RS256"}
|
||||
with pytest.raises(Exception):
|
||||
await ext._get_signing_key("fake-token")
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Authentication — JWKS mode
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestAuthenticateJWKS:
|
||||
"""Tests for JWKS-based JWT verification."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_valid_token(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock()
|
||||
ext._context = mock_context
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": VALID_UUID, "aud": "authenticated"}
|
||||
|
||||
result = await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
assert isinstance(result, TenantContext)
|
||||
expected_schema = "user_" + VALID_UUID.replace("-", "_")
|
||||
assert result.schema_name == expected_schema
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_custom_prefix(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
ext.schema_prefix = "org"
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock()
|
||||
ext._context = mock_context
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": VALID_UUID}
|
||||
|
||||
result = await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
assert result.schema_name.startswith("org_")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_expired_token(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch(
|
||||
"hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode",
|
||||
side_effect=pyjwt.ExpiredSignatureError(),
|
||||
),
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Token has expired"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_invalid_audience(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch(
|
||||
"hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode",
|
||||
side_effect=pyjwt.InvalidAudienceError(),
|
||||
),
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Invalid token audience"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_invalid_issuer(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch(
|
||||
"hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode",
|
||||
side_effect=pyjwt.InvalidIssuerError(),
|
||||
),
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Invalid token issuer"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_decode_error(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch(
|
||||
"hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode",
|
||||
side_effect=pyjwt.DecodeError(),
|
||||
),
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Invalid token"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_missing_sub_claim(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"email": "test@example.com"} # no sub
|
||||
|
||||
with pytest.raises(AuthenticationError, match="missing subject"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_empty_sub_claim(self):
|
||||
"""Empty string sub claim should be treated as missing."""
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": ""}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="missing subject"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_generic_exception(self):
|
||||
"""Unexpected exceptions during decode should be caught and wrapped."""
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch(
|
||||
"hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode",
|
||||
side_effect=RuntimeError("unexpected internal error"),
|
||||
),
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Token verification failed"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Authentication — Legacy mode
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestAuthenticateLegacy:
|
||||
"""Tests for legacy /auth/v1/user endpoint verification."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_valid_token(self):
|
||||
ext, mock_client = _setup_legacy_ext()
|
||||
mock_client.get.return_value = _make_mock_response(200, {"id": VALID_UUID})
|
||||
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock()
|
||||
ext._context = mock_context
|
||||
|
||||
result = await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
assert isinstance(result, TenantContext)
|
||||
expected_schema = "user_" + VALID_UUID.replace("-", "_")
|
||||
assert result.schema_name == expected_schema
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_calls_user_endpoint(self):
|
||||
ext, mock_client = _setup_legacy_ext()
|
||||
mock_client.get.return_value = _make_mock_response(200, {"id": VALID_UUID})
|
||||
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock()
|
||||
ext._context = mock_context
|
||||
|
||||
token = _make_valid_token()
|
||||
await ext.authenticate(RequestContext(api_key=token))
|
||||
|
||||
mock_client.get.assert_called_once_with(
|
||||
"https://test.supabase.co/auth/v1/user",
|
||||
headers={
|
||||
"Authorization": f"Bearer {token}",
|
||||
"apikey": "test-service-key",
|
||||
},
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_expired_token_401(self):
|
||||
ext, mock_client = _setup_legacy_ext()
|
||||
mock_client.get.return_value = _make_mock_response(401)
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Invalid or expired token"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_supabase_error_500(self):
|
||||
ext, mock_client = _setup_legacy_ext()
|
||||
mock_client.get.return_value = _make_mock_response(500)
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Authentication failed: 500"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_no_user_id(self):
|
||||
ext, mock_client = _setup_legacy_ext()
|
||||
mock_client.get.return_value = _make_mock_response(200, {"email": "test@example.com"})
|
||||
|
||||
with pytest.raises(AuthenticationError, match="no user ID found"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_timeout(self):
|
||||
ext, mock_client = _setup_legacy_ext()
|
||||
mock_client.get.side_effect = httpx.TimeoutException("Request timed out")
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Authentication timeout"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_connection_error(self):
|
||||
ext, mock_client = _setup_legacy_ext()
|
||||
mock_client.get.side_effect = httpx.ConnectError("Connection refused")
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Connection error"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Authentication — common (both modes)
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestAuthenticateCommon:
|
||||
"""Tests that apply regardless of verification mode."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_missing_token(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Missing Authorization header"):
|
||||
await ext.authenticate(RequestContext(api_key=None))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_empty_token(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Missing Authorization header"):
|
||||
await ext.authenticate(RequestContext(api_key=""))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_short_token(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Invalid token format"):
|
||||
await ext.authenticate(RequestContext(api_key="short"))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_not_initialized(self):
|
||||
ext = _make_extension()
|
||||
# _http_client is None by default
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Extension not initialized"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_rejects_non_uuid_user_id(self):
|
||||
"""User IDs that aren't valid UUIDs should be rejected for schema safety."""
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": "not-a-uuid"}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Invalid user ID format"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_rejects_malicious_user_id(self):
|
||||
"""User IDs with SQL injection attempts should be rejected."""
|
||||
ext, _ = _setup_jwks_ext()
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": "'; DROP TABLE users;--"}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Invalid user ID format"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Schema management
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestSupabaseTenantExtensionSchemaManagement:
|
||||
"""Tests for schema initialization and caching."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schema_initialized_on_first_access(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock()
|
||||
ext._context = mock_context
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": VALID_UUID}
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
expected_schema = "user_" + VALID_UUID.replace("-", "_")
|
||||
mock_context.run_migration.assert_called_once_with(expected_schema)
|
||||
assert expected_schema in ext._initialized_schemas
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schema_cached_on_second_access(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock()
|
||||
ext._context = mock_context
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": VALID_UUID}
|
||||
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
# run_migration should only be called once
|
||||
expected_schema = "user_" + VALID_UUID.replace("-", "_")
|
||||
mock_context.run_migration.assert_called_once_with(expected_schema)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schema_init_failure(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock(side_effect=RuntimeError("Migration failed"))
|
||||
ext._context = mock_context
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": VALID_UUID}
|
||||
|
||||
with pytest.raises(AuthenticationError, match="Failed to initialize tenant"):
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
# Schema should NOT be cached on failure
|
||||
expected_schema = "user_" + VALID_UUID.replace("-", "_")
|
||||
assert expected_schema not in ext._initialized_schemas
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# List tenants
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestSupabaseTenantExtensionListTenants:
|
||||
"""Tests for list_tenants behavior."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tenants_empty(self):
|
||||
ext = _make_extension()
|
||||
tenants = await ext.list_tenants()
|
||||
assert tenants == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tenants_after_auth(self):
|
||||
ext, _ = _setup_jwks_ext()
|
||||
mock_context = AsyncMock(spec=ExtensionContext)
|
||||
mock_context.run_migration = AsyncMock()
|
||||
ext._context = mock_context
|
||||
|
||||
with (
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.get_unverified_header") as mock_header,
|
||||
patch("hindsight_api.extensions.builtin.supabase_tenant.pyjwt.decode") as mock_decode,
|
||||
):
|
||||
mock_header.return_value = {"kid": "test-key-1", "alg": "RS256"}
|
||||
mock_decode.return_value = {"sub": VALID_UUID}
|
||||
await ext.authenticate(RequestContext(api_key=_make_valid_token()))
|
||||
|
||||
tenants = await ext.list_tenants()
|
||||
assert len(tenants) == 1
|
||||
assert isinstance(tenants[0], Tenant)
|
||||
expected_schema = "user_" + VALID_UUID.replace("-", "_")
|
||||
assert tenants[0].schema == expected_schema
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Shutdown
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestSupabaseTenantExtensionShutdown:
|
||||
"""Tests for on_shutdown behavior."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_shutdown_closes_client(self):
|
||||
ext = _make_extension()
|
||||
mock_client = AsyncMock(spec=httpx.AsyncClient)
|
||||
ext._http_client = mock_client
|
||||
|
||||
await ext.on_shutdown()
|
||||
|
||||
mock_client.aclose.assert_called_once()
|
||||
assert ext._http_client is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_shutdown_no_client(self):
|
||||
ext = _make_extension()
|
||||
# _http_client is None by default — should not raise
|
||||
await ext.on_shutdown()
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Extension loader integration
|
||||
# ======================================================================
|
||||
|
||||
|
||||
class TestSupabaseTenantExtensionLoader:
|
||||
"""Tests for loading via the extension loader."""
|
||||
|
||||
def test_load_via_extension_loader(self, monkeypatch):
|
||||
monkeypatch.setenv(
|
||||
"HINDSIGHT_API_TENANT_EXTENSION",
|
||||
"hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension",
|
||||
)
|
||||
monkeypatch.setenv("HINDSIGHT_API_TENANT_SUPABASE_URL", "https://test.supabase.co")
|
||||
monkeypatch.setenv("HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY", "test-key")
|
||||
monkeypatch.setenv("HINDSIGHT_API_TENANT_SCHEMA_PREFIX", "custom")
|
||||
|
||||
ext = load_extension("TENANT", TenantExtension)
|
||||
|
||||
assert ext is not None
|
||||
assert isinstance(ext, SupabaseTenantExtension)
|
||||
assert ext.supabase_url == "https://test.supabase.co"
|
||||
assert ext.supabase_service_key == "test-key"
|
||||
assert ext.schema_prefix == "custom"
|
||||
|
||||
def test_load_without_service_key(self, monkeypatch):
|
||||
"""Extension should load without service key — JWKS mode doesn't need it."""
|
||||
monkeypatch.setenv(
|
||||
"HINDSIGHT_API_TENANT_EXTENSION",
|
||||
"hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension",
|
||||
)
|
||||
monkeypatch.setenv("HINDSIGHT_API_TENANT_SUPABASE_URL", "https://test.supabase.co")
|
||||
monkeypatch.delenv("HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY", raising=False)
|
||||
|
||||
ext = load_extension("TENANT", TenantExtension)
|
||||
|
||||
assert ext is not None
|
||||
assert isinstance(ext, SupabaseTenantExtension)
|
||||
assert ext.supabase_service_key is None
|
||||
147
hindsight-integrations/supabase/README.md
Normal file
147
hindsight-integrations/supabase/README.md
Normal file
|
|
@ -0,0 +1,147 @@
|
|||
# Supabase Tenant Extension for Hindsight
|
||||
|
||||
A built-in TenantExtension that validates [Supabase](https://supabase.com) JWTs and provides multi-tenant memory isolation. Each authenticated user gets their own PostgreSQL schema, ensuring complete data separation.
|
||||
|
||||
## Features
|
||||
|
||||
- **Local JWT Verification** - Validates tokens locally using JWKS public keys (no network call per request)
|
||||
- **Automatic Schema Isolation** - Each user gets `{prefix}_{user_id}` schema
|
||||
- **Zero User Management** - Leverages your existing Supabase Auth setup
|
||||
- **Production Ready** - Includes health checks, timeouts, key rotation handling, and error handling
|
||||
- **Built-in** - Ships with Hindsight, no extra installation needed
|
||||
- **Legacy Support** - Falls back to `/auth/v1/user` endpoint for HS256 projects
|
||||
|
||||
## Configuration
|
||||
|
||||
The Supabase tenant extension is built into Hindsight. Just set the environment variables:
|
||||
|
||||
```bash
|
||||
# Required
|
||||
HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension
|
||||
HINDSIGHT_API_TENANT_SUPABASE_URL=https://your-project.supabase.co
|
||||
|
||||
# Optional - only needed for legacy HS256 projects or startup health check
|
||||
HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY=your-service-role-key
|
||||
|
||||
# Optional
|
||||
HINDSIGHT_API_TENANT_SCHEMA_PREFIX=user # Default: "user"
|
||||
```
|
||||
|
||||
> **Note:** Most Supabase projects use asymmetric JWT signing (ES256/RS256) and the extension verifies tokens locally using JWKS — no service key needed. The `service_role` key is only required if your project uses legacy HS256 signing or if you want the startup health check.
|
||||
|
||||
## Usage
|
||||
|
||||
Clients pass their Supabase access token in the Authorization header:
|
||||
|
||||
```bash
|
||||
# Get user's access token from Supabase Auth
|
||||
TOKEN=$(curl -s -X POST "https://your-project.supabase.co/auth/v1/token?grant_type=password" \
|
||||
-H "apikey: your-anon-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"email": "user@example.com", "password": "xxx"}' | jq -r '.access_token')
|
||||
|
||||
# Use with Hindsight
|
||||
curl -X POST "https://your-hindsight-server/v1/default/banks/my-bank/memories" \
|
||||
-H "Authorization: Bearer $TOKEN" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"items": [{"content": "User preference: likes dark mode"}]}'
|
||||
```
|
||||
|
||||
## How It Works
|
||||
|
||||
```
|
||||
┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐
|
||||
│ Your App │ │ Hindsight │ │ Supabase │
|
||||
│ │ │ │ │ │
|
||||
│ 1. User logs │ │ │ │ JWKS keys │
|
||||
│ in via │────▶│ │ │ fetched once │
|
||||
│ Supabase │ │ │ │ on startup │
|
||||
│ │ │ │ │ │
|
||||
│ 2. App calls │ │ 3. Extension │ │ │
|
||||
│ Hindsight │────▶│ verifies │ │ │
|
||||
│ with JWT │ │ JWT locally │ │ │
|
||||
│ │ │ (no network │ │ │
|
||||
│ │ │ call) │ │ │
|
||||
│ │ │ │ │ │
|
||||
│ │ │ 4. Routes to │ │ │
|
||||
│ │◀────│ user's │ │ │
|
||||
│ │ │ schema │ │ │
|
||||
└─────────────────┘ └─────────────────┘ └─────────────────┘
|
||||
```
|
||||
|
||||
1. On startup, Hindsight fetches JWKS public keys from Supabase (cached for 10 minutes)
|
||||
2. User authenticates with your app via Supabase Auth
|
||||
3. Your app calls Hindsight API with the user's JWT
|
||||
4. Extension verifies the JWT signature locally using cached public keys
|
||||
5. On success, routes request to user's isolated schema (`user_{uuid}`)
|
||||
|
||||
For legacy HS256 projects, the extension falls back to calling `/auth/v1/user` per request.
|
||||
|
||||
## Schema Isolation
|
||||
|
||||
Each user gets a completely isolated PostgreSQL schema:
|
||||
|
||||
```
|
||||
Hindsight Database
|
||||
├── Schema: user_abc123_def456 (User A)
|
||||
│ ├── memories
|
||||
│ ├── entities
|
||||
│ └── ...
|
||||
├── Schema: user_xyz789_... (User B)
|
||||
│ ├── memories
|
||||
│ ├── entities
|
||||
│ └── ...
|
||||
└── Schema: public (Hindsight internals)
|
||||
```
|
||||
|
||||
User A cannot access User B's data - they're in separate schemas.
|
||||
|
||||
## Configuration Options
|
||||
|
||||
| Variable | Required | Default | Description |
|
||||
|----------|----------|---------|-------------|
|
||||
| `HINDSIGHT_API_TENANT_SUPABASE_URL` | Yes | - | Your Supabase project URL |
|
||||
| `HINDSIGHT_API_TENANT_SUPABASE_SERVICE_KEY` | No | - | Supabase service_role key (only needed for HS256 projects or health check) |
|
||||
| `HINDSIGHT_API_TENANT_SCHEMA_PREFIX` | No | `user` | Prefix for schema names (must be a valid Postgres identifier) |
|
||||
|
||||
## Deployment Examples
|
||||
|
||||
### Docker
|
||||
|
||||
```dockerfile
|
||||
FROM ghcr.io/vectorize-io/hindsight:latest
|
||||
|
||||
ENV HINDSIGHT_API_TENANT_EXTENSION=hindsight_api.extensions.builtin.supabase_tenant:SupabaseTenantExtension
|
||||
```
|
||||
|
||||
### Railway
|
||||
|
||||
```toml
|
||||
# railway.toml
|
||||
[build]
|
||||
builder = "dockerfile"
|
||||
dockerfilePath = "Dockerfile"
|
||||
|
||||
[deploy]
|
||||
healthcheckPath = "/health"
|
||||
```
|
||||
|
||||
Set environment variables in Railway dashboard.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
| Error | Cause | Solution |
|
||||
|-------|-------|----------|
|
||||
| `401 Unauthorized` | Invalid or expired JWT | Get fresh token from Supabase |
|
||||
| `Missing Authorization header` | No Bearer token sent | Add `Authorization: Bearer <token>` header |
|
||||
| `Unable to find signing key` | JWT signed with unknown key | Check Supabase JWT algorithm settings |
|
||||
| `Authentication timeout` | Supabase slow/unreachable (legacy mode) | Check Supabase status, retry |
|
||||
| `SUPABASE_SERVICE_KEY is required when JWKS is not available` | HS256 project without service key | Provide service_role key or switch to asymmetric JWT signing |
|
||||
|
||||
## Contributing
|
||||
|
||||
This extension was originally developed by [BrighterBalance](https://brighterbalance.app) for their AI advisor product.
|
||||
|
||||
## License
|
||||
|
||||
MIT
|
||||
Loading…
Reference in a new issue