diff --git a/docs/framework/multi-tenancy.md b/docs/framework/multi-tenancy.md index 9ce0f2cb..82675570 100644 --- a/docs/framework/multi-tenancy.md +++ b/docs/framework/multi-tenancy.md @@ -137,6 +137,25 @@ then the session's choice validated against a membership (#363). Without one it `tenant_id` claim, and for **anonymous** requests only, the configured `tenant_header`. An authenticated user can never pick a tenant by header. +### Source and `Vary` + +The middleware records where the tenant came from on +`request.state.tenant_source`: `fixed` (`default_tenant`), `subdomain`, +`header` (a member's explicit per-request choice), `session`, `claim` (the +principal's `tenant_id`), `anon_header` (anonymous requests only), +`resolver` (a custom resolver that returned a bare string), or `None` when no +tenant was bound. Use it to tell a deliberate per-request choice from an +ambient session default. + +A resolver may return `str | None` (source `resolver`), a +`(tenant_id, source)` pair, or `TenantResolution(tenant_id, source, vary)` from +`simple_module_hosting.middleware`. `vary` lists the request headers the answer +depended on; the middleware merges them into the response `Vary` (existing +entries are kept, duplicates are dropped case-insensitively, `Vary: *` is left +alone). The `tenants` resolver reports `Host` when subdomains are enabled and +the tenant header whenever it is configured, even if the answer was `None`, so +a shared cache cannot serve one tenant's response to another. + ## Unique keys On a tenant-scoped table every business key is per tenant: put `tenant_id` in diff --git a/framework/hosting/simple_module_hosting/_tenant.py b/framework/hosting/simple_module_hosting/_tenant.py index 1f6296a0..0513bd36 100644 --- a/framework/hosting/simple_module_hosting/_tenant.py +++ b/framework/hosting/simple_module_hosting/_tenant.py @@ -2,10 +2,12 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, MutableMapping +from dataclasses import dataclass +from typing import Any from simple_module_db import current_tenant_id, is_valid_tenant_id -from starlette.datastructures import Headers +from starlette.datastructures import Headers, MutableHeaders from starlette.requests import Request from starlette.types import ASGIApp, Receive, Scope, Send @@ -14,13 +16,56 @@ TENANT_HEADER = "X-Tenant-ID" -TenantResolver = Callable[[Request], Awaitable[str | None]] +@dataclass(frozen=True, slots=True) +class TenantResolution: + """A resolver's answer plus how it got there. + + ``source`` is what ends up on ``request.state.tenant_source`` (``"subdomain"``, + ``"header"``, ``"session"``, ...). ``vary`` names the request headers the + answer depended on (``"Host"``, the tenant header); the middleware adds them + to the response ``Vary`` so a shared cache cannot hand one tenant's response + to another. Report a header even when the answer is ``None``: its absence + was an input too. + """ + + tenant_id: str | None + source: str | None = None + vary: tuple[str, ...] = () + + +TenantResolver = Callable[[Request], Awaitable["str | TenantResolution | tuple | None"]] """Module-owned tenant resolution, registered as ``app.state.tenant_resolver``. Returns the tenant the request acts for, or ``None``. It owns *every* source — membership, session, subdomain, header — so it is also where each is validated. Without one the middleware falls back to the principal's -``tenant_id`` claim.""" +``tenant_id`` claim. + +A plain ``str | None`` still works (the source is recorded as ``"resolver"``). +Return a :class:`TenantResolution` (or a ``(tenant_id, source)`` pair) to +report the source and the headers consulted.""" + + +def _normalise(result: object) -> TenantResolution: + if isinstance(result, TenantResolution): + return result + if isinstance(result, tuple): + tenant_id, source = result[0], (result[1] if len(result) > 1 else None) + return TenantResolution(tenant_id, source if tenant_id is not None else None) + return TenantResolution(result, "resolver" if result is not None else None) # type: ignore[arg-type] + + +def merge_vary(existing: str | None, names: tuple[str, ...]) -> str | None: + """``existing`` with ``names`` appended, case-insensitively de-duplicated.""" + present = [v.strip() for v in (existing or "").split(",") if v.strip()] + if "*" in present: + return existing + seen = {v.lower() for v in present} + for name in names: + if name.lower() not in seen: + present.append(name) + seen.add(name.lower()) + return ", ".join(present) if present else None class TenantMiddleware: @@ -30,7 +75,10 @@ class TenantMiddleware: :class:`~simple_module_db.mixins.MultiTenantMixin` models are automatically filtered, and new objects get ``tenant_id`` populated. - Also stores the resolved value on ``request.state.tenant_id``. + Also stores the resolved value on ``request.state.tenant_id`` and where it + came from on ``request.state.tenant_source`` (``fixed``, ``subdomain``, + ``header``, ``session``, ``claim``, ``anon_header``, ``resolver`` or ``None``). + Headers the answer depended on are added to the response ``Vary``. Resolution: @@ -59,9 +107,12 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: return request = Request(scope) - tenant_id = await self._resolve(request, scope) + resolution = await self._resolve(request, scope) + tenant_id = resolution.tenant_id request.state.tenant_id = tenant_id + request.state.tenant_source = resolution.source + send = self._vary_sender(send, resolution.vary) if tenant_id is not None: token = current_tenant_id.set(tenant_id) try: @@ -72,26 +123,51 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: await self.app(scope, receive, send) - async def _resolve(self, request: Request, scope: Scope) -> str | None: + @staticmethod + def _vary_sender(send: Send, vary: tuple[str, ...]) -> Send: + if not vary: + return send + + async def wrapped(message: MutableMapping[str, Any]) -> None: + if message["type"] == "http.response.start": + headers = MutableHeaders(scope=message) + merged = merge_vary(headers.get("vary"), vary) + if merged is not None: + headers["vary"] = merged + await send(message) + + return wrapped + + async def _resolve(self, request: Request, scope: Scope) -> TenantResolution: if self.fixed is not None: - return self.fixed + return TenantResolution(self.fixed, "fixed") app = scope.get("app") resolver: TenantResolver | None = getattr( getattr(app, "state", None), "tenant_resolver", None ) if resolver is not None: - return await resolver(request) + return _normalise(await resolver(request)) user = getattr(request.state, "user", None) if user is not None: - return getattr(user, "tenant_id", None) + claim = getattr(user, "tenant_id", None) + return TenantResolution(claim, "claim" if claim is not None else None) if self.header: value = Headers(scope=scope).get(self.header) # Unvalidated, an over-long value is a 500 on the first stamped # write (VARCHAR(50)) and any junk becomes a tenant name (#366). - return value if is_valid_tenant_id(value) else None - return None - - -__all__ = ["TENANT_HEADER", "TenantMiddleware", "TenantResolver"] + ok = is_valid_tenant_id(value) + return TenantResolution( + value if ok else None, "anon_header" if ok else None, (self.header,) + ) + return TenantResolution(None) + + +__all__ = [ + "TENANT_HEADER", + "TenantMiddleware", + "TenantResolution", + "TenantResolver", + "merge_vary", +] diff --git a/framework/hosting/simple_module_hosting/middleware.py b/framework/hosting/simple_module_hosting/middleware.py index f7184d94..c75018a7 100644 --- a/framework/hosting/simple_module_hosting/middleware.py +++ b/framework/hosting/simple_module_hosting/middleware.py @@ -26,7 +26,12 @@ CorrelationIdMiddleware, RequestLoggingMiddleware, ) -from simple_module_hosting._tenant import TENANT_HEADER, TenantMiddleware, TenantResolver +from simple_module_hosting._tenant import ( + TENANT_HEADER, + TenantMiddleware, + TenantResolution, + TenantResolver, +) from simple_module_hosting.permissions import expand_permissions, resolve_permissions if TYPE_CHECKING: @@ -62,6 +67,7 @@ "RequestLoggingMiddleware", "SecurityHeadersMiddleware", "TenantMiddleware", + "TenantResolution", "TenantResolver", ] diff --git a/framework/hosting/tests/test_tenant_source.py b/framework/hosting/tests/test_tenant_source.py new file mode 100644 index 00000000..983cc53d --- /dev/null +++ b/framework/hosting/tests/test_tenant_source.py @@ -0,0 +1,112 @@ +"""TenantMiddleware records where the tenant came from and merges ``Vary``.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +from simple_module_hosting.middleware import TenantMiddleware, TenantResolution + + +def _scope(headers=None, user=None, resolver=None): + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": headers or [], + "state": {}, + "app": SimpleNamespace(state=SimpleNamespace(tenant_resolver=resolver)), + } + if user is not None: + scope["state"]["user"] = user + return scope + + +async def _receive(): # pragma: no cover + return {"type": "http.request", "body": b"", "more_body": False} + + +async def _run(mw_kwargs, scope, vary_header: bytes | None = None): + sent: list[dict] = [] + + async def inner(scope, receive, send): + start = {"type": "http.response.start", "status": 200, "headers": []} + if vary_header is not None: + start["headers"].append((b"vary", vary_header)) + await send(start) + + async def send(message): + sent.append(message) + + await TenantMiddleware(inner, **mw_kwargs)(scope, _receive, send) + headers = [v.decode() for k, v in sent[0]["headers"] if k == b"vary"] + return scope["state"], headers + + +async def test_fixed_source(): + state, vary = await _run({"fixed": "main"}, _scope()) + assert (state["tenant_id"], state["tenant_source"]) == ("main", "fixed") + assert vary == [] + + +async def test_claim_source(): + state, _ = await _run({}, _scope(user=SimpleNamespace(tenant_id="acme"))) + assert state["tenant_source"] == "claim" + + +async def test_claim_absent_has_no_source(): + state, _ = await _run({}, _scope(user=SimpleNamespace(tenant_id=None))) + assert state["tenant_source"] is None + + +async def test_anon_header_source_and_vary(): + scope = _scope(headers=[(b"x-tenant-id", b"acme")]) + state, vary = await _run({"header": "X-Tenant-ID"}, scope) + assert state["tenant_source"] == "anon_header" + assert vary == ["X-Tenant-ID"] + + +async def test_anon_header_consulted_but_invalid_still_varies(): + scope = _scope(headers=[(b"x-tenant-id", b"x" * 80)]) + state, vary = await _run({"header": "X-Tenant-ID"}, scope) + assert (state["tenant_id"], state["tenant_source"]) == (None, None) + assert vary == ["X-Tenant-ID"] + + +async def test_no_source_no_vary(): + state, vary = await _run({}, _scope()) + assert state["tenant_source"] is None + assert vary == [] + + +@pytest.mark.parametrize( + ("answer", "expected"), + [ + ("acme", ("acme", "resolver")), + (None, (None, None)), + (("acme", "session"), ("acme", "session")), + (TenantResolution("acme", "subdomain"), ("acme", "subdomain")), + ], +) +async def test_resolver_return_shapes(answer, expected): + async def resolver(request): + return answer + + state, _ = await _run({}, _scope(resolver=resolver)) + assert (state["tenant_id"], state["tenant_source"]) == expected + + +async def test_resolver_vary_is_merged_not_clobbered(): + async def resolver(request): + return TenantResolution("acme", "header", ("Host", "X-Tenant-ID")) + + _, vary = await _run({}, _scope(resolver=resolver), vary_header=b"Accept-Encoding, host") + assert vary == ["Accept-Encoding, host, X-Tenant-ID"] # Host already present + + +async def test_vary_star_is_left_alone(): + async def resolver(request): + return TenantResolution("acme", "header", ("X-Tenant-ID",)) + + _, vary = await _run({}, _scope(resolver=resolver), vary_header=b"*") + assert vary == ["*"] diff --git a/modules/tenants/tenants/host_resolver.py b/modules/tenants/tenants/host_resolver.py index 80986a70..cf13f360 100644 --- a/modules/tenants/tenants/host_resolver.py +++ b/modules/tenants/tenants/host_resolver.py @@ -36,6 +36,10 @@ def forget_hosts() -> None: _BY_SLUG.clear() +def subdomains_enabled(request: Request) -> bool: + return bool(request.app.state.tenants.settings.subdomain_base.strip().strip(".")) + + def subdomain_slug(request: Request) -> str | None: """The tenant slug named by the request's host, if subdomains are enabled.""" base = request.app.state.tenants.settings.subdomain_base.strip().lower().strip(".") diff --git a/modules/tenants/tenants/resolver.py b/modules/tenants/tenants/resolver.py index e391907d..f219de3c 100644 --- a/modules/tenants/tenants/resolver.py +++ b/modules/tenants/tenants/resolver.py @@ -22,6 +22,7 @@ from fastapi import FastAPI from simple_module_core.invalidation import Invalidation, InvalidationBus from simple_module_db import is_valid_tenant_id +from simple_module_hosting.middleware import TenantResolution from starlette.requests import Request from tenants.constants import ( @@ -36,6 +37,7 @@ is_public_route, resolve_from_host, subdomain_slug, + subdomains_enabled, ) from tenants.service import TenantService @@ -124,34 +126,48 @@ def _with_tenant_role(user: Any, tenant_id: str, role: str) -> Any: return dataclasses.replace(user, **changes) -async def resolve_tenant(request: Request) -> str | None: - """``TenantResolver`` for the framework's ``TenantMiddleware``.""" +async def resolve_tenant(request: Request) -> TenantResolution: + """``TenantResolver`` for the framework's ``TenantMiddleware``. + + Reports the source (``subdomain`` / ``header`` / ``session``) and the headers + the answer depended on, so the middleware can set ``Vary``. + """ + tenant_id, source, vary = await _resolve(request) + return TenantResolution(tenant_id, source if tenant_id is not None else None, vary) + + +async def _resolve(request: Request) -> tuple[str | None, str | None, tuple[str, ...]]: request.state.tenant_role = None request.state.tenant_suspended = False request.state.suspended_tenant_name = None user = getattr(request.state, "user", None) slug = subdomain_slug(request) if slug is not None: - return await _resolve_subdomain(request, user, slug) + return await _resolve_subdomain(request, user, slug), "subdomain", ("Host",) + # No slug in the host is still an answer that depended on the host. + vary: tuple[str, ...] = ("Host",) if subdomains_enabled(request) else () if user is None: - return None + return None, None, vary user_id = str(user.id) memberships = await memberships_for(request.app, user_id) - requested = _header_tenant(request) + header_name = _header_name(request) + if header_name: + vary = (*vary, header_name) + requested = request.headers.get(header_name) if header_name else None if requested is not None: # An explicit per-request choice (API clients). Never fall back to # another tenant: a client that asked for X — or sent junk — must not # act on Y. if not is_valid_tenant_id(requested): - return None + return None, None, vary active = next( (m for m in memberships if m.id == requested and m.status == TenantStatus.ACTIVE), None, ) if active is None: - return None - return _enter(request, user, active) + return None, None, vary + return _enter(request, user, active), "header", vary session = request.scope.get("session") preferred = session.get(SESSION_ACTIVE_TENANT) if session is not None else None @@ -169,10 +185,10 @@ async def resolve_tenant(request: Request) -> str | None: request.state.tenant_suspended = any( m.status == TenantStatus.SUSPENDED for m in memberships ) - return None + return None, None, vary if session is not None and preferred != active.id and not chosen_suspended: session[SESSION_ACTIVE_TENANT] = active.id - return _enter(request, user, active) + return _enter(request, user, active), "session", vary async def _resolve_subdomain(request: Request, user: Any, slug: str) -> str | None: @@ -189,10 +205,9 @@ async def _resolve_subdomain(request: Request, user: Any, slug: str) -> str | No return tenant_id -def _header_tenant(request: Request) -> str | None: +def _header_name(request: Request) -> str: settings = getattr(getattr(request.app.state, "sm", None), "settings", None) - header = getattr(settings, "tenant_header", "") or "" - return request.headers.get(header) if header else None + return getattr(settings, "tenant_header", "") or "" def _enter(request: Request, user: Any, active: MyTenantView) -> str: diff --git a/modules/tenants/tests/test_resolver_source.py b/modules/tenants/tests/test_resolver_source.py new file mode 100644 index 00000000..9500a7e3 --- /dev/null +++ b/modules/tenants/tests/test_resolver_source.py @@ -0,0 +1,77 @@ +"""``request.state.tenant_source`` and ``Vary`` from the tenants resolver.""" + +from __future__ import annotations + +import httpx +import pytest +from fastapi import Request +from tenants.host_resolver import forget_hosts + + +@pytest.fixture +async def probe(app): + async def route(request: Request): + s = request.state + return { + "tenant": getattr(s, "tenant_id", None), + "source": getattr(s, "tenant_source", None), + } + + app.add_api_route("/probe", route, methods=["GET"]) + app.state.public_routes.add_prefix("/probe", methods={"GET"}) + yield app + app.state.tenants.settings.subdomain_base = "" + forget_hosts() + + +def _vary(resp) -> list[str]: + return [t.strip().lower() for t in resp.headers.get("vary", "").split(",") if t.strip()] + + +async def _create(client, name): + return (await client.post("/api/tenants/", json={"name": name})).json() + + +async def test_session_source(probe, user_client): + async with user_client("a@x.io") as (a, _): + t = await _create(a, "One") + body = (await a.get("/probe")).json() + assert body == {"tenant": t["id"], "source": "session"} + + +async def test_header_source_and_vary(probe, user_client): + async with user_client("a@x.io") as (a, _): + t = await _create(a, "One") + resp = await a.get("/probe", headers={"X-Tenant-ID": t["id"]}) + assert resp.json()["source"] == "header" + assert "x-tenant-id" in _vary(resp) + + +async def test_unresolved_header_has_no_source_but_varies(probe, user_client): + async with user_client("a@x.io") as (a, _): + await _create(a, "One") + resp = await a.get("/probe", headers={"X-Tenant-ID": "nope"}) + assert resp.json()["source"] is None + assert "x-tenant-id" in _vary(resp) + + +async def test_subdomain_source_and_vary_host(probe, user_client): + probe.state.tenants.settings.subdomain_base = "example.com" + forget_hosts() + async with user_client("a@x.io") as (a, _): + t = await _create(a, "Acme") + anon = httpx.AsyncClient( + transport=httpx.ASGITransport(app=probe), base_url=f"http://{t['slug']}.example.com" + ) + async with anon: + resp = await anon.get("/probe") + assert resp.json() == {"tenant": t["id"], "source": "subdomain"} + assert "host" in _vary(resp) + + +async def test_vary_has_no_duplicates(probe, user_client): + async with user_client("a@x.io") as (a, _): + await _create(a, "One") + resp = await a.get("/probe", headers={"X-Tenant-ID": "nope"}) + tokens = _vary(resp) + assert len(tokens) == len(set(tokens))