376 lines
13 KiB
Python
376 lines
13 KiB
Python
"""Bearer-token verification (Phase 4.2 — DB-backed with in-memory cache).
|
|
|
|
The store is fed by two sources:
|
|
|
|
1. **DB** (``Token`` table from Phase 4.1) — the primary source. Rows
|
|
carry Argon2id hashes; the admin UI (Phase 4.3) will be the only
|
|
writer at steady state.
|
|
2. **`BOOTSTRAP_TOKENS`** env — a name→plaintext map kept around for
|
|
first-run, disaster-recovery, and ``examples/`` smoke tests. These
|
|
entries live alongside DB rows in the cache and are never persisted.
|
|
|
|
Hot path is in-memory: at :meth:`start` we pull every non-revoked DB
|
|
row and stash it in a list; subsequent :meth:`verify` calls re-load
|
|
when the cache is older than ``ttl_seconds``. ``last_used_at`` updates
|
|
are coalesced into a small dict and flushed by a background task every
|
|
``flush_interval`` seconds — one transaction per flush rather than one
|
|
per request.
|
|
|
|
We can't index DB rows by a derived plaintext key because Argon2 salts
|
|
are random — so verify does a linear scan over candidates, calling
|
|
``argon2.PasswordHasher.verify`` on each. N is small by design (single
|
|
operator, ~10 tokens at most); the cost is irrelevant. The scan runs
|
|
through ``asyncio.to_thread`` to keep the event loop free of the ~50ms
|
|
KDF block.
|
|
|
|
The module knows nothing about HTTP frameworks. It takes a raw token
|
|
(or a verbatim ``Authorization`` header value) and returns a
|
|
:class:`TokenIdentity` (name + scope + db-id), or ``None`` for a miss.
|
|
Frontends own the 401 response shape.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import contextlib
|
|
import hmac
|
|
import logging
|
|
import time
|
|
from dataclasses import dataclass
|
|
from datetime import UTC, datetime
|
|
from typing import TYPE_CHECKING
|
|
|
|
from argon2 import PasswordHasher
|
|
from argon2.exceptions import InvalidHashError, VerifyMismatchError
|
|
|
|
from beaver_gateway.storage import list_active_tokens, touch_token
|
|
|
|
if TYPE_CHECKING:
|
|
from collections.abc import Mapping
|
|
|
|
from beaver_gateway.storage import Database
|
|
|
|
|
|
_log = logging.getLogger("beaver_gateway.auth")
|
|
|
|
_BOOTSTRAP_SCOPE = "*"
|
|
|
|
|
|
class TokenStoreError(ValueError):
|
|
"""Malformed ``BOOTSTRAP_TOKENS`` value or duplicate token."""
|
|
|
|
|
|
VALID_SCOPES: frozenset[str] = frozenset({"*", "messages", "mcp", "admin", "api"})
|
|
"""The scopes a ``Token.scope`` may hold (Phase 4.3 admin UI enforces).
|
|
|
|
* ``*`` — wildcard, may use any frontend
|
|
* ``messages`` — Anthropic Messages frontend only
|
|
* ``mcp`` — MCP server frontend only
|
|
* ``api`` — conversations API (``/api``) and its SSE
|
|
* ``admin`` — reserved for programmatic admin access; the AdminFrontend
|
|
itself authenticates via session cookies, not bearer tokens, so this
|
|
scope is unused today and kept for forward compatibility.
|
|
"""
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class TokenIdentity:
|
|
"""What :meth:`TokenStore.verify` resolves to on success.
|
|
|
|
``token_id`` is the DB row id for persisted tokens, or ``None`` for
|
|
an env-bootstrap match (those have no DB row to touch). ``scope``
|
|
gates which frontend the token may hit (see :data:`VALID_SCOPES`);
|
|
bootstrap tokens implicitly get ``"*"``.
|
|
"""
|
|
|
|
name: str
|
|
scope: str
|
|
token_id: int | None
|
|
|
|
def allows(self, required: str) -> bool:
|
|
"""``True`` when this identity may access a route gated by ``required``.
|
|
|
|
``"*"`` is the wildcard; an exact match satisfies a single scope.
|
|
Unknown ``required`` values intentionally fall through to a
|
|
strict equality check — callers should pass one of
|
|
:data:`VALID_SCOPES`.
|
|
"""
|
|
return self.scope in ("*", required)
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _CachedToken:
|
|
"""One non-revoked row, copied out of the DB into the hot-path cache."""
|
|
|
|
id: int
|
|
name: str
|
|
scope: str
|
|
hashed_value: str
|
|
|
|
|
|
# Default Argon2id parameters from ``argon2-cffi`` are fine for our scope.
|
|
# They target ~50ms on a modern CPU — enough to make a stolen-hash brute
|
|
# force expensive, cheap enough to verify a handful per request.
|
|
_HASHER = PasswordHasher()
|
|
|
|
|
|
def hash_token(plaintext: str) -> str:
|
|
"""Return an Argon2id hash for ``plaintext`` (admin / seed-only path).
|
|
|
|
Phase 4.3 will call this when the admin creates a token; Phase 4.2
|
|
exposes it so smoke scripts can seed the DB without re-implementing
|
|
the same line.
|
|
"""
|
|
return _HASHER.hash(plaintext)
|
|
|
|
|
|
class TokenStore:
|
|
"""DB-backed verifier with in-memory cache + TTL + batched touches.
|
|
|
|
Construct in ``cli.main`` after :class:`Database` is up, then
|
|
``await store.start()`` to prime the cache and spin up the flusher
|
|
task. ``await store.stop()`` on shutdown drains the touch queue.
|
|
|
|
Bootstrap entries (from ``BOOTSTRAP_TOKENS``) sit alongside DB rows
|
|
in the same lookup path; we check them first, in constant time, so
|
|
they remain usable even if the DB is unreachable. They never appear
|
|
in ``last_used_at`` flushes because they have no DB row.
|
|
"""
|
|
|
|
__slots__ = (
|
|
"_bootstrap_by_value",
|
|
"_bootstrap_scopes",
|
|
"_cache",
|
|
"_db",
|
|
"_flush_interval",
|
|
"_flusher_task",
|
|
"_loaded_at",
|
|
"_lock",
|
|
"_touch_queue",
|
|
"_ttl",
|
|
)
|
|
|
|
def __init__(
|
|
self,
|
|
db: Database | None = None,
|
|
*,
|
|
bootstrap: Mapping[str, str] | None = None,
|
|
ttl_seconds: float = 30.0,
|
|
flush_interval: float = 5.0,
|
|
) -> None:
|
|
# Bootstrap is keyed by value internally so verify is O(1) over
|
|
# plaintext. Each value also keeps its name for audit lines.
|
|
by_value: dict[str, str] = {}
|
|
for name, value in (bootstrap or {}).items():
|
|
if not name or not value:
|
|
msg = f"empty name or value in bootstrap map (name={name!r})"
|
|
raise TokenStoreError(msg)
|
|
if value in by_value:
|
|
msg = (
|
|
f"duplicate bootstrap token value for names "
|
|
f"{by_value[value]!r} and {name!r}"
|
|
)
|
|
raise TokenStoreError(msg)
|
|
by_value[value] = name
|
|
self._bootstrap_by_value: dict[str, str] = by_value
|
|
self._bootstrap_scopes: dict[str, str] = dict.fromkeys(
|
|
by_value.values(), _BOOTSTRAP_SCOPE
|
|
)
|
|
|
|
self._db = db
|
|
self._ttl = ttl_seconds
|
|
self._flush_interval = flush_interval
|
|
|
|
self._cache: list[_CachedToken] = []
|
|
self._loaded_at: float = 0.0
|
|
self._lock = asyncio.Lock()
|
|
self._touch_queue: dict[int, datetime] = {}
|
|
self._flusher_task: asyncio.Task[None] | None = None
|
|
|
|
# ---- bootstrap parsing (kept for `cli` / tests) ---------------------
|
|
|
|
@staticmethod
|
|
def parse_bootstrap(raw: str) -> dict[str, str]:
|
|
"""Parse ``name1:value1,name2:value2`` (the ``BOOTSTRAP_TOKENS`` form)."""
|
|
tokens: dict[str, str] = {}
|
|
for chunk in raw.split(","):
|
|
entry = chunk.strip()
|
|
if not entry:
|
|
continue
|
|
name, sep, value = entry.partition(":")
|
|
if not sep:
|
|
msg = f"token entry missing ':' separator: {entry!r}"
|
|
raise TokenStoreError(msg)
|
|
name, value = name.strip(), value.strip()
|
|
if name in tokens:
|
|
msg = f"duplicate token name: {name!r}"
|
|
raise TokenStoreError(msg)
|
|
tokens[name] = value
|
|
return tokens
|
|
|
|
@classmethod
|
|
def from_env(cls, raw: str, db: Database | None = None) -> TokenStore:
|
|
"""Legacy entrypoint: bootstrap-only (or bootstrap + db).
|
|
|
|
Phase 1.3 call sites still expect a one-liner; we keep the
|
|
classmethod so they don't have to learn the new constructor.
|
|
"""
|
|
return cls(db, bootstrap=cls.parse_bootstrap(raw))
|
|
|
|
# ---- lifecycle ------------------------------------------------------
|
|
|
|
async def start(self) -> None:
|
|
"""Prime the cache and (if a DB is attached) start the flusher loop."""
|
|
await self._refresh()
|
|
if self._db is not None:
|
|
self._flusher_task = asyncio.create_task(
|
|
self._flusher_loop(), name="beaver-gateway.token-flusher"
|
|
)
|
|
|
|
async def stop(self) -> None:
|
|
"""Cancel the flusher and run one final drain."""
|
|
if self._flusher_task is not None:
|
|
self._flusher_task.cancel()
|
|
with contextlib.suppress(asyncio.CancelledError):
|
|
await self._flusher_task
|
|
self._flusher_task = None
|
|
await self._flush_now()
|
|
|
|
async def invalidate(self) -> None:
|
|
"""Force the next verify to re-read from DB (Phase 4.3 admin hook)."""
|
|
self._loaded_at = 0.0
|
|
|
|
# ---- verify path ----------------------------------------------------
|
|
|
|
async def verify(self, token: str | None) -> TokenIdentity | None:
|
|
"""Return the matching identity, or ``None`` for unknown/empty tokens."""
|
|
if not token:
|
|
return None
|
|
|
|
# Bootstrap first: constant-time compare per entry, never hits DB.
|
|
# `compare_digest` is overkill for a name→value lookup but cheap
|
|
# and removes one timing variable for free.
|
|
for value, name in self._bootstrap_by_value.items():
|
|
if hmac.compare_digest(token, value):
|
|
return TokenIdentity(
|
|
name=name,
|
|
scope=self._bootstrap_scopes.get(name, _BOOTSTRAP_SCOPE),
|
|
token_id=None,
|
|
)
|
|
|
|
if self._db is None:
|
|
return None
|
|
|
|
await self._ensure_fresh()
|
|
|
|
# Snapshot the cache reference so a refresh mid-scan doesn't
|
|
# surprise us. List itself is immutable per refresh (we swap,
|
|
# not mutate).
|
|
cache = self._cache
|
|
for entry in cache:
|
|
try:
|
|
await asyncio.to_thread(_HASHER.verify, entry.hashed_value, token)
|
|
except VerifyMismatchError:
|
|
continue
|
|
except InvalidHashError:
|
|
_log.warning(
|
|
"token row %d has an unparseable hash — skipping", entry.id
|
|
)
|
|
continue
|
|
self._touch_queue[entry.id] = datetime.now(UTC)
|
|
return TokenIdentity(name=entry.name, scope=entry.scope, token_id=entry.id)
|
|
return None
|
|
|
|
async def verify_bearer(self, authorization: str | None) -> TokenIdentity | None:
|
|
"""Strip the ``Bearer`` prefix (case-insensitive) then verify.
|
|
|
|
Accepts a bare token too — Cursor's MCP transport sometimes
|
|
passes the raw value via ``?token=`` and reuses the same
|
|
verifier; treating an unprefixed header as a bare token keeps
|
|
both call sites on one method.
|
|
"""
|
|
if not authorization:
|
|
return None
|
|
head, sep, rest = authorization.partition(" ")
|
|
token = rest.strip() if sep and head.lower() == "bearer" else authorization
|
|
return await self.verify(token)
|
|
|
|
def __len__(self) -> int:
|
|
return len(self._cache) + len(self._bootstrap_by_value)
|
|
|
|
def __bool__(self) -> bool:
|
|
return bool(self._cache) or bool(self._bootstrap_by_value)
|
|
|
|
# ---- internals ------------------------------------------------------
|
|
|
|
async def _ensure_fresh(self) -> None:
|
|
if self._db is None:
|
|
return
|
|
now = time.monotonic()
|
|
if now - self._loaded_at <= self._ttl:
|
|
return
|
|
async with self._lock:
|
|
# Re-check under the lock — first arrival reloaded, others
|
|
# should fall through.
|
|
now = time.monotonic()
|
|
if now - self._loaded_at <= self._ttl:
|
|
return
|
|
await self._refresh()
|
|
|
|
async def _refresh(self) -> None:
|
|
if self._db is None:
|
|
self._loaded_at = time.monotonic()
|
|
return
|
|
async with self._db.session() as session:
|
|
rows = await list_active_tokens(session)
|
|
next_cache: list[_CachedToken] = []
|
|
for row in rows:
|
|
if row.id is None:
|
|
# Defensive: SQLModel will assign an id on insert; a
|
|
# None here would mean someone handed us an unsaved row.
|
|
continue
|
|
next_cache.append(
|
|
_CachedToken(
|
|
id=row.id,
|
|
name=row.name,
|
|
scope=row.scope,
|
|
hashed_value=row.hashed_value,
|
|
)
|
|
)
|
|
self._cache = next_cache
|
|
self._loaded_at = time.monotonic()
|
|
_log.debug("token cache refreshed: %d active row(s)", len(next_cache))
|
|
|
|
async def _flusher_loop(self) -> None:
|
|
try:
|
|
while True:
|
|
await asyncio.sleep(self._flush_interval)
|
|
await self._flush_now()
|
|
except asyncio.CancelledError:
|
|
raise
|
|
except Exception: # noqa: BLE001 — never let the flusher die silently
|
|
_log.exception("token flusher crashed; touches will stop")
|
|
|
|
async def _flush_now(self) -> None:
|
|
if self._db is None or not self._touch_queue:
|
|
return
|
|
# Detach the queue so concurrent verify() writes don't bleed
|
|
# into the in-flight transaction.
|
|
pending, self._touch_queue = self._touch_queue, {}
|
|
async with self._db.session() as session:
|
|
for token_id in pending:
|
|
# We don't pass the timestamp through — `touch_token`
|
|
# stamps `now` itself, and we'd rather have one source
|
|
# of truth than reconcile clocks.
|
|
await touch_token(session, token_id=token_id)
|
|
_log.debug("flushed %d token touch(es)", len(pending))
|
|
|
|
|
|
__all__ = [
|
|
"VALID_SCOPES",
|
|
"TokenIdentity",
|
|
"TokenStore",
|
|
"TokenStoreError",
|
|
"hash_token",
|
|
]
|