Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions docs/framework/multi-tenancy.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
106 changes: 91 additions & 15 deletions framework/hosting/simple_module_hosting/_tenant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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:
Expand All @@ -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:

Expand Down Expand Up @@ -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:
Expand All @@ -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",
]
8 changes: 7 additions & 1 deletion framework/hosting/simple_module_hosting/middleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -62,6 +67,7 @@
"RequestLoggingMiddleware",
"SecurityHeadersMiddleware",
"TenantMiddleware",
"TenantResolution",
"TenantResolver",
]

Expand Down
112 changes: 112 additions & 0 deletions framework/hosting/tests/test_tenant_source.py
Original file line number Diff line number Diff line change
@@ -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 == ["*"]
4 changes: 4 additions & 0 deletions modules/tenants/tenants/host_resolver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(".")
Expand Down
Loading
Loading