diff --git a/app/api/messages.py b/app/api/messages.py index e1b6750..b15a107 100644 --- a/app/api/messages.py +++ b/app/api/messages.py @@ -536,7 +536,7 @@ def _end_trace_spans(error=None): ) # Record usage - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -632,7 +632,7 @@ def _end_trace_spans(error=None): ) # Record usage - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -716,7 +716,7 @@ async def _web_search_stream_with_usage(): ws_error = str(e) raise finally: - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -746,7 +746,7 @@ async def _web_search_stream_with_usage(): ) # Record usage - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -824,7 +824,7 @@ async def _web_fetch_stream_with_usage(): wf_error = str(e) raise finally: - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -854,7 +854,7 @@ async def _web_fetch_stream_with_usage(): ) # Record usage - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -959,7 +959,7 @@ async def _web_fetch_stream_with_usage(): else: result = await provider.invoke(request_data, target_model, api_key_info) # Record usage - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=target_model, @@ -1060,7 +1060,7 @@ async def _web_fetch_stream_with_usage(): _turn_span.end() # Record usage - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -1090,7 +1090,7 @@ async def _web_fetch_stream_with_usage(): print(f"[ERROR] HTTP Status: {e.http_status}") print(f"[ERROR] Error Type: {e.error_type}\n") - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -1117,7 +1117,7 @@ async def _web_fetch_stream_with_usage(): exc_info=True, ) - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, @@ -1379,7 +1379,7 @@ async def _handle_streaming_request( finally: # Record usage after stream completes - usage_tracker.record_usage( + usage_tracker.record_usage_nowait( api_key=api_key_info.get("api_key"), request_id=request_id, model=request_data.model, diff --git a/app/api/openai_passthrough/router.py b/app/api/openai_passthrough/router.py index aba84e4..5eab654 100644 --- a/app/api/openai_passthrough/router.py +++ b/app/api/openai_passthrough/router.py @@ -204,7 +204,7 @@ def _record_usage( _, usage, _ = _managers() norm = normalize_usage(raw_usage, api_surface) try: - usage.record_usage( + usage.record_usage_nowait( api_key=api_key_info.get("api_key", ""), request_id=str(uuid4()), model=model, diff --git a/app/core/config.py b/app/core/config.py index 8b23735..ebb4243 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -124,6 +124,15 @@ class Settings(BaseSettings): api_key_header: str = Field(default="x-api-key", alias="API_KEY_HEADER") require_api_key: bool = Field(default=True, alias="REQUIRE_API_KEY") master_api_key: Optional[str] = Field(default=None, alias="MASTER_API_KEY") + api_key_cache_ttl_seconds: int = Field( + default=60, alias="API_KEY_CACHE_TTL_SECONDS", + description=( + "TTL in seconds for the in-process API key validation cache " + "(0 to disable). Avoids a DynamoDB read per request; key " + "changes (create/disable) made in another process take up to " + "this long to apply on running workers." + ) + ) # Rate Limiting Settings rate_limit_enabled: bool = Field(default=True, alias="RATE_LIMIT_ENABLED") @@ -290,6 +299,15 @@ class Settings(BaseSettings): "inference profile ARNs to their underlying foundation model ID.", ) + # Model Mapping Cache + model_mapping_cache_ttl_seconds: int = Field( + default=300, + alias="MODEL_MAPPING_CACHE_TTL_SECONDS", + description="TTL (seconds) for the in-process model mapping cache " + "(0 to disable). Mapping changes made in another process " + "(e.g. the admin portal) take up to this long to apply.", + ) + # Beta features that require InvokeModel API instead of Converse API # These features are only available via InvokeModel/InvokeModelWithResponseStream beta_headers_requiring_invoke_model: List[str] = Field( diff --git a/app/core/metrics.py b/app/core/metrics.py index 1619dae..d825ddd 100644 --- a/app/core/metrics.py +++ b/app/core/metrics.py @@ -3,11 +3,10 @@ Provides Prometheus-compatible metrics for monitoring. """ -from prometheus_client import Counter, Histogram, Gauge, Info -from typing import Optional -from app.core.config import settings +from prometheus_client import Counter, Gauge, Histogram, Info +from app.core.config import settings # Request metrics request_counter = Counter( @@ -69,6 +68,12 @@ ["api_key"], ) +# Usage tracking metrics +usage_writes_dropped_counter = Counter( + "usage_writes_dropped_total", + "Usage/billing rows dropped because the background write backlog was full", +) + # Authentication metrics auth_failures_counter = Counter( "auth_failures_total", diff --git a/app/core/ttl_cache.py b/app/core/ttl_cache.py new file mode 100644 index 0000000..383daca --- /dev/null +++ b/app/core/ttl_cache.py @@ -0,0 +1,71 @@ +"""Minimal thread-safe TTL cache shared by hot-path lookups. + +Used to keep per-request DynamoDB reads (API key validation, model +mapping) off the request path. Deliberately small instead of pulling in +cachetools: get/set/invalidate with per-entry expiry and a hard entry +bound so client-controlled keys cannot grow memory without limit. +""" +import threading +import time +from typing import Any + + +class TTLCache: + """Thread-safe key/value cache with per-entry TTL and a size bound. + + ``get`` returns ``(hit, value)`` so a cached ``None`` (negative + caching) is distinguishable from a miss. + + Eviction: when full, expired entries are purged first; if the cache + is still full, the soonest-expiring entries are dropped. Cache keys + are client-controlled input, and negative entries carry the shortest + TTLs — so spam evicts itself before it can evict hot positive + entries. Entries are cheap to recompute (single DynamoDB reads), so + bounded memory matters more than hit-rate precision. + """ + + def __init__(self, max_entries: int = 10_000): + self._max_entries = max_entries + self._entries: dict[Any, tuple[Any, float]] = {} + self._lock = threading.Lock() + + def get(self, key: Any) -> tuple[bool, Any]: + with self._lock: + entry = self._entries.get(key) + if entry is None: + return False, None + value, expires_at = entry + if time.monotonic() >= expires_at: + del self._entries[key] + return False, None + return True, value + + def set(self, key: Any, value: Any, ttl_seconds: float) -> None: + with self._lock: + if key not in self._entries and len(self._entries) >= self._max_entries: + self._evict_locked() + self._entries[key] = (value, time.monotonic() + ttl_seconds) + + def invalidate(self, key: Any) -> None: + with self._lock: + self._entries.pop(key, None) + + def clear(self) -> None: + with self._lock: + self._entries.clear() + + def __len__(self) -> int: + with self._lock: + return len(self._entries) + + def _evict_locked(self) -> None: + """Purge expired entries; drop the soonest-expiring if still full.""" + now = time.monotonic() + expired = [k for k, (_, exp) in self._entries.items() if now >= exp] + for k in expired: + del self._entries[k] + if len(self._entries) >= self._max_entries: + n_drop = max(1, len(self._entries) // 10) + by_expiry = sorted(self._entries.items(), key=lambda kv: kv[1][1]) + for k, _ in by_expiry[:n_drop]: + del self._entries[k] diff --git a/app/db/dynamodb.py b/app/db/dynamodb.py index dc3c7f6..bbc362d 100644 --- a/app/db/dynamodb.py +++ b/app/db/dynamodb.py @@ -4,8 +4,12 @@ Provides interfaces for interacting with DynamoDB tables for API keys, usage tracking, and model mapping. """ +import inspect import json +import logging +import threading import time +from concurrent.futures import Future, ThreadPoolExecutor from datetime import datetime, timedelta, timezone from decimal import Decimal from typing import Any, Dict, List, Optional, Union @@ -15,8 +19,67 @@ from botocore.exceptions import ClientError from app.core.config import settings +from app.core.ttl_cache import TTLCache from app.services.inference_profile_resolver import get_inference_profile_resolver +logger = logging.getLogger(__name__) + +# Single background thread for usage writes: keeps the synchronous +# put_item off the event loop and serializes writes (one writer thread; +# readers elsewhere share the underlying thread-safe boto3 client). +_usage_write_executor: Optional[ThreadPoolExecutor] = None +_usage_write_executor_lock = threading.Lock() + +# Backpressure bound: if DynamoDB hangs while traffic continues, queued +# writes (and their kwargs) would otherwise accumulate without limit. +# Past this depth new writes are dropped (counted + logged) — bounded +# memory is worth more than a lossless backlog that OOMs the proxy. +_MAX_PENDING_USAGE_WRITES = 10_000 +_pending_usage_writes = 0 +_pending_usage_writes_lock = threading.Lock() + + +def _get_usage_write_executor() -> ThreadPoolExecutor: + """Lazily create the shared usage-writer executor (thread-safe).""" + global _usage_write_executor + if _usage_write_executor is None: + with _usage_write_executor_lock: + if _usage_write_executor is None: + _usage_write_executor = ThreadPoolExecutor( + max_workers=1, thread_name_prefix="usage-writer" + ) + return _usage_write_executor + + +def _completed_future() -> Future: + f: Future = Future() + f.set_result(None) + return f + + +def drain_usage_writes(timeout: float = 5.0) -> int: + """Best-effort flush of queued usage writes (for app shutdown). + + Blocks until the backlog is empty or the deadline passes. Returns + the number of writes still pending at the deadline (0 = fully + drained); pending writes are logged as lost. + """ + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + with _pending_usage_writes_lock: + if _pending_usage_writes == 0: + return 0 + time.sleep(0.02) + with _pending_usage_writes_lock: + remaining = _pending_usage_writes + if remaining: + logger.warning( + "Shutdown deadline reached with %d usage writes still pending; " + "those usage/billing rows are lost", + remaining, + ) + return remaining + class DynamoDBClient: """DynamoDB client for managing tables and operations.""" @@ -448,36 +511,41 @@ def validate_api_key(self, api_key: str) -> Optional[Dict[str, Any]]: api_key: API key to validate Returns: - API key details if valid, None otherwise + API key details if valid, None if the key does not exist or + is inactive. + + Raises: + ClientError: On DynamoDB failures (throttling, 5xx). Callers + must treat this as "lookup failed", not "invalid key" — + the auth middleware negative-caches None results, so + conflating the two would lock valid keys out during a + DynamoDB throttling event. """ - try: - response = self.table.get_item(Key={"api_key": api_key}) - item = response.get("Item") + response = self.table.get_item(Key={"api_key": api_key}) + item = response.get("Item") - if not item: - return None + if not item: + return None - # If key is active, return it - if item.get("is_active", False): - return item + # If key is active, return it + if item.get("is_active", False): + return item - # Check if key was deactivated due to budget exceeded - # and if a new month has started - auto-reactivate it - deactivated_reason = item.get("deactivated_reason") - if deactivated_reason == "budget_exceeded": - budget_mtd_month = item.get("budget_mtd_month", "") - current_month = datetime.now(timezone.utc).strftime("%Y-%m") + # Check if key was deactivated due to budget exceeded + # and if a new month has started - auto-reactivate it + deactivated_reason = item.get("deactivated_reason") + if deactivated_reason == "budget_exceeded": + budget_mtd_month = item.get("budget_mtd_month", "") + current_month = datetime.now(timezone.utc).strftime("%Y-%m") - if budget_mtd_month != current_month: - # New month has started - reactivate and reset MTD - self._reactivate_for_new_month(api_key, current_month) - # Fetch updated item - response = self.table.get_item(Key={"api_key": api_key}) - return response.get("Item") + if budget_mtd_month != current_month: + # New month has started - reactivate and reset MTD + self._reactivate_for_new_month(api_key, current_month) + # Fetch updated item + response = self.table.get_item(Key={"api_key": api_key}) + return response.get("Item") - return None - except ClientError: - return None + return None def _reactivate_for_new_month(self, api_key: str, current_month: str) -> bool: """ @@ -1015,6 +1083,7 @@ def record_usage( cache_ttl: Optional[str] = None, api_surface: Optional[str] = None, reasoning_tokens: int = 0, + timestamp_ms: Optional[int] = None, ): """ Record API usage. @@ -1033,10 +1102,16 @@ def record_usage( cache_ttl: Effective cache TTL used ("5m" or "1h"), for billing differentiation api_surface: Source endpoint family ("messages", "chat_completions", or "responses") reasoning_tokens: Reasoning tokens (already counted in output_tokens; stored separately for visibility) + timestamp_ms: Event time in epoch milliseconds. Defaults to now. + The usage table key is (api_key, timestamp), so millisecond + precision keeps same-second rows from overwriting each other, + and callers that queue writes should stamp at submit time. """ # Use string timestamp to match CDK table schema (STRING type) current_time = int(time.time()) - timestamp = str(current_time * 1000) # milliseconds as string + if timestamp_ms is None: + timestamp_ms = int(time.time() * 1000) + timestamp = str(timestamp_ms) resolved_model = _safe_resolve_model(model) metadata = dict(metadata) if metadata else {} @@ -1073,6 +1148,75 @@ def record_usage( self.table.put_item(Item=item) + def record_usage_nowait(self, **kwargs) -> Future: + """Record API usage on a background thread without blocking the caller. + + Request handlers run on the event loop; the synchronous + ``put_item`` in :meth:`record_usage` would otherwise freeze every + in-flight request (including streaming chunks) for the duration + of the DynamoDB round trip. + + Never raises past argument validation: write failures, a full + backlog, and a shut-down executor are logged and swallowed — + call sites sit inside except/finally blocks, and a usage-write + error must never mask the original exception or fail a request. + Bad kwargs DO raise TypeError at the call site (a typo must not + become silently dropped billing data). The event timestamp is + stamped here, at submit time, so rows keep their true time even + if the backlog drains late. + + Returns the Future so tests (or shutdown hooks) can wait on it. + Use :func:`drain_usage_writes` to flush on app shutdown. + """ + global _pending_usage_writes + + inspect.signature(self.record_usage).bind(**kwargs) + kwargs.setdefault("timestamp_ms", int(time.time() * 1000)) + + with _pending_usage_writes_lock: + if _pending_usage_writes >= _MAX_PENDING_USAGE_WRITES: + backlog = _pending_usage_writes + else: + backlog = None + _pending_usage_writes += 1 + if backlog is not None: + logger.warning( + "Usage write dropped: backlog full (%d pending) — " + "usage/billing row lost for request_id=%s", + backlog, + kwargs.get("request_id"), + ) + try: + from app.core.metrics import usage_writes_dropped_counter + + usage_writes_dropped_counter.inc() + except Exception: + pass + return _completed_future() + + def _write(): + global _pending_usage_writes + try: + self.record_usage(**kwargs) + except Exception: + logger.warning( + "Background usage write failed for request_id=%s", + kwargs.get("request_id"), + exc_info=True, + ) + finally: + with _pending_usage_writes_lock: + _pending_usage_writes -= 1 + + try: + return _get_usage_write_executor().submit(_write) + except RuntimeError as e: + # Executor already shut down (app/interpreter exit) + with _pending_usage_writes_lock: + _pending_usage_writes -= 1 + logger.warning("Usage write skipped, writer shut down: %s", e) + return _completed_future() + def get_usage_stats( self, api_key: str, start_time: Optional[datetime] = None, end_time: Optional[datetime] = None ) -> Dict[str, Any]: @@ -1139,30 +1283,53 @@ def get_usage_stats( class ModelMappingManager: """Manager for custom model mappings.""" + # Shared across instances: converters build a fresh manager per + # request, so an instance-level cache would never see a hit. Cached + # None entries matter too — clients passing Bedrock ARNs directly + # (pass-through) would otherwise pay a DynamoDB read per request. + _cache = TTLCache() + def __init__(self, dynamodb_client: DynamoDBClient): """Initialize model mapping manager.""" self.dynamodb = dynamodb_client.dynamodb self.table = self.dynamodb.Table(dynamodb_client.model_mapping_table_name) + self._cache_ttl = settings.model_mapping_cache_ttl_seconds def get_mapping(self, anthropic_model_id: str) -> Optional[str]: """ Get Bedrock model ID for an Anthropic model ID. + Results (including "no mapping") are cached for + ``settings.model_mapping_cache_ttl_seconds``; mapping changes + made in another process (e.g. the admin portal) take up to that + long to apply here. + Args: anthropic_model_id: Anthropic model identifier Returns: Bedrock model ARN or None """ + if self._cache_ttl > 0: + hit, cached = self._cache.get(anthropic_model_id) + if hit: + return cached + try: response = self.table.get_item( Key={"anthropic_model_id": anthropic_model_id} ) - item = response.get("Item") - return item.get("bedrock_model_id") if item else None except ClientError: + # Transient failure: fall back to default mapping upstream, + # but never cache it — that would pin None for the full TTL. return None + item = response.get("Item") + mapping = item.get("bedrock_model_id") if item else None + if self._cache_ttl > 0: + self._cache.set(anthropic_model_id, mapping, self._cache_ttl) + return mapping + def set_mapping(self, anthropic_model_id: str, bedrock_model_id: str): """ Set custom model mapping. @@ -1177,6 +1344,10 @@ def set_mapping(self, anthropic_model_id: str, bedrock_model_id: str): "updated_at": int(time.time()), } self.table.put_item(Item=item) + if self._cache_ttl > 0: + self._cache.set(anthropic_model_id, bedrock_model_id, self._cache_ttl) + else: + self._cache.invalidate(anthropic_model_id) def delete_mapping(self, anthropic_model_id: str): """ @@ -1186,6 +1357,7 @@ def delete_mapping(self, anthropic_model_id: str): anthropic_model_id: Anthropic model identifier """ self.table.delete_item(Key={"anthropic_model_id": anthropic_model_id}) + self._cache.invalidate(anthropic_model_id) def list_mappings(self) -> List[Dict[str, str]]: """ diff --git a/app/main.py b/app/main.py index c589001..7632736 100644 --- a/app/main.py +++ b/app/main.py @@ -212,6 +212,10 @@ async def lifespan(app: FastAPI): # Shutdown print("Shutting down application...") + # Flush queued usage/billing writes with a bounded deadline + from app.db.dynamodb import drain_usage_writes + await asyncio.to_thread(drain_usage_writes, 5.0) + # Shutdown tracing if settings.enable_tracing: from app.tracing import shutdown_tracing diff --git a/app/middleware/auth.py b/app/middleware/auth.py index 722ea9b..42c7962 100644 --- a/app/middleware/auth.py +++ b/app/middleware/auth.py @@ -2,17 +2,52 @@ Authentication middleware for API key validation. Validates API keys from request headers and attaches user information to requests. + +Validation results are cached in-process with a TTL so the hot path does +not pay a DynamoDB round trip per request, and cache misses run in a +dedicated worker thread pool so the synchronous boto3 call never blocks +the event loop (and never queues behind long-running default-executor +work like web search or docker pulls). """ +import asyncio +import copy +import hashlib import hmac -from typing import Callable +import threading +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor from fastapi import HTTPException, Request, status from fastapi.security import APIKeyHeader from starlette.middleware.base import BaseHTTPMiddleware from app.core.config import settings +from app.core.ttl_cache import TTLCache from app.db.dynamodb import APIKeyManager, DynamoDBClient +# Invalid keys are cached only briefly: long enough to blunt brute-force +# spam against DynamoDB, short enough that a freshly created key works +# almost immediately. +NEGATIVE_CACHE_TTL_SECONDS = 5.0 + +# Small dedicated pool for auth lookups. The event loop's default +# executor is shared with multi-second work elsewhere in this codebase +# (Tavily web search/fetch, docker pulls); auth is on the critical path +# of every request and must not queue behind those. +_auth_executor: ThreadPoolExecutor | None = None +_auth_executor_lock = threading.Lock() + + +def _get_auth_executor() -> ThreadPoolExecutor: + global _auth_executor + if _auth_executor is None: + with _auth_executor_lock: + if _auth_executor is None: + _auth_executor = ThreadPoolExecutor( + max_workers=4, thread_name_prefix="auth-validate" + ) + return _auth_executor + # API Key header scheme api_key_header_scheme = APIKeyHeader( @@ -24,16 +59,91 @@ class AuthMiddleware(BaseHTTPMiddleware): """Middleware for API key authentication.""" - def __init__(self, app, dynamodb_client: DynamoDBClient): + def __init__( + self, + app, + dynamodb_client: DynamoDBClient, + cache_ttl_seconds: float | None = None, + ): """ Initialize auth middleware. Args: app: FastAPI application dynamodb_client: DynamoDB client instance + cache_ttl_seconds: TTL for cached validation results. + Defaults to ``settings.api_key_cache_ttl_seconds``; + 0 disables caching. Key changes made in another process + (e.g. the admin portal) take up to this long to apply. """ super().__init__(app) self.api_key_manager = APIKeyManager(dynamodb_client) + if cache_ttl_seconds is None: + cache_ttl_seconds = settings.api_key_cache_ttl_seconds + self._cache_ttl = cache_ttl_seconds + self._cache = TTLCache() + # In-flight lookups by cache key (single-flight): only touched + # from the event loop, so no lock is needed. + self._inflight: dict[str, asyncio.Task] = {} + + async def _validate_api_key(self, api_key: str) -> dict | None: + """Validate a key via cache, falling back to DynamoDB off-loop. + + The cache is keyed by a SHA-256 digest of the key, so plaintext + credentials are not held as dict keys and attacker-supplied + oversized "keys" cannot inflate per-entry memory. Values are + deep-copied in and out so handlers mutating their + ``api_key_info`` (including nested dicts) cannot poison the + cache. Concurrent misses for the same key coalesce into one + DynamoDB read (single-flight). Validation errors are treated as + invalid but never cached: a transient DynamoDB failure must not + lock a good key out for the TTL window. + """ + cache_key = hashlib.sha256(api_key.encode()).hexdigest() + if self._cache_ttl > 0: + hit, cached = self._cache.get(cache_key) + if hit: + return copy.deepcopy(cached) if cached is not None else None + + task = self._inflight.get(cache_key) + is_leader = task is None + if task is None: + loop = asyncio.get_running_loop() + task = loop.create_task(self._lookup(api_key)) + self._inflight[cache_key] = task + + try: + api_key_info = await task + except Exception as e: + if is_leader: + print("\n[ERROR] Exception during API key validation") + print(f"[ERROR] Type: {type(e).__name__}") + print(f"[ERROR] Message: {str(e)}") + import traceback + print(f"[ERROR] Traceback:\n{traceback.format_exc()}\n") + return None + finally: + if is_leader: + self._inflight.pop(cache_key, None) + + if is_leader and self._cache_ttl > 0: + if api_key_info: + self._cache.set( + cache_key, copy.deepcopy(api_key_info), self._cache_ttl + ) + else: + self._cache.set( + cache_key, + None, + min(self._cache_ttl, NEGATIVE_CACHE_TTL_SECONDS), + ) + return copy.deepcopy(api_key_info) if api_key_info is not None else None + + async def _lookup(self, api_key: str) -> dict | None: + loop = asyncio.get_running_loop() + return await loop.run_in_executor( + _get_auth_executor(), self.api_key_manager.validate_api_key, api_key + ) async def dispatch(self, request: Request, call_next: Callable): """ @@ -91,16 +201,8 @@ async def dispatch(self, request: Request, call_next: Callable): } return await call_next(request) - # Validate API key in DynamoDB - try: - api_key_info = self.api_key_manager.validate_api_key(api_key) - except Exception as e: - print(f"\n[ERROR] Exception during API key validation") - print(f"[ERROR] Type: {type(e).__name__}") - print(f"[ERROR] Message: {str(e)}") - import traceback - print(f"[ERROR] Traceback:\n{traceback.format_exc()}\n") - api_key_info = None + # Validate API key (in-process cache, DynamoDB on miss) + api_key_info = await self._validate_api_key(api_key) if not api_key_info: # Deliberately do not log the rejected key (or any derivative of diff --git a/env.example b/env.example index 27d263c..28b799e 100644 --- a/env.example +++ b/env.example @@ -38,6 +38,13 @@ USAGE_TTL_DAYS=30 API_KEY_HEADER=x-api-key REQUIRE_API_KEY=True MASTER_API_KEY=sk-your-master-key +# In-process API key validation cache TTL in seconds (0 to disable). +# Key changes made in the admin portal take up to this long to apply. +API_KEY_CACHE_TTL_SECONDS=60 + +# In-process model mapping cache TTL in seconds (0 to disable). +# Mapping changes made in the admin portal take up to this long to apply. +MODEL_MAPPING_CACHE_TTL_SECONDS=300 # Rate Limiting RATE_LIMIT_ENABLED=True diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..32c1c49 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,16 @@ +"""Shared test fixtures. + +Process-wide caches must not leak state between tests: without this, +tests that exercise the real managers become order-dependent (a mapping +cached by one test silently satisfies or breaks a later one). +""" +import pytest + + +@pytest.fixture(autouse=True) +def _clear_shared_caches(): + from app.db.dynamodb import ModelMappingManager + + ModelMappingManager._cache.clear() + yield + ModelMappingManager._cache.clear() diff --git a/tests/integration/test_openai_passthrough/conftest.py b/tests/integration/test_openai_passthrough/conftest.py index 80ef0b1..db393a8 100644 --- a/tests/integration/test_openai_passthrough/conftest.py +++ b/tests/integration/test_openai_passthrough/conftest.py @@ -55,7 +55,10 @@ def mock_model_mapping_manager(): @pytest.fixture def mock_usage_tracker(): - tracker = MagicMock() + from app.db.dynamodb import UsageTracker + + # spec= so assertions fail if the real method is renamed or removed + tracker = MagicMock(spec=UsageTracker) with patch("app.api.openai_passthrough.router.UsageTracker", return_value=tracker): yield tracker diff --git a/tests/integration/test_openai_passthrough/test_chat_completions.py b/tests/integration/test_openai_passthrough/test_chat_completions.py index 6430449..450bf1b 100644 --- a/tests/integration/test_openai_passthrough/test_chat_completions.py +++ b/tests/integration/test_openai_passthrough/test_chat_completions.py @@ -62,8 +62,8 @@ def test_non_streaming_chat_completions_forwards_and_logs_usage( assert sent_body["input"] == [{"role": "user", "content": "hi"}] assert "store" not in sent_body # Usage was recorded - assert mock_usage_tracker.record_usage.called - kwargs = mock_usage_tracker.record_usage.call_args.kwargs + assert mock_usage_tracker.record_usage_nowait.called + kwargs = mock_usage_tracker.record_usage_nowait.call_args.kwargs assert kwargs["input_tokens"] == 10 assert kwargs["output_tokens"] == 5 assert kwargs["cached_tokens"] == 3 @@ -458,7 +458,7 @@ def test_upstream_4xx_returned_verbatim(client, respx_mock, mock_usage_tracker): ) assert r.status_code == 404 assert r.json() == err_body - assert not mock_usage_tracker.record_usage.called # Don't log usage on errors + assert not mock_usage_tracker.record_usage_nowait.called # Don't log usage on errors def test_missing_auth_returns_401(client): @@ -506,8 +506,8 @@ def test_streaming_chat_completions_forwards_sse_and_records_usage( assert b'"delta":{"content":"hi"}' in out assert b"[DONE]" in out # Usage recorded from the chunk that had it - assert mock_usage_tracker.record_usage.called - kw = mock_usage_tracker.record_usage.call_args.kwargs + assert mock_usage_tracker.record_usage_nowait.called + kw = mock_usage_tracker.record_usage_nowait.call_args.kwargs assert kw["input_tokens"] == 7 assert kw["output_tokens"] == 2 assert kw["cached_tokens"] == 1 @@ -613,7 +613,7 @@ def test_streaming_chat_completions_without_include_usage_does_not_log( ) as r: list(r.iter_bytes()) # drain - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_streaming_chat_completions_does_not_inject_event_lines( @@ -697,4 +697,4 @@ def test_streaming_upstream_timeout_returns_json_504( body = r.json() assert body["error"]["type"] == "upstream_error" assert "timeout" in body["error"]["message"].lower() - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called diff --git a/tests/integration/test_openai_passthrough/test_responses.py b/tests/integration/test_openai_passthrough/test_responses.py index ca8a251..d202e22 100644 --- a/tests/integration/test_openai_passthrough/test_responses.py +++ b/tests/integration/test_openai_passthrough/test_responses.py @@ -53,7 +53,7 @@ def test_non_streaming_responses_forwards_and_logs_usage( assert r.status_code == 200 assert r.json() == upstream assert route.called - kw = mock_usage_tracker.record_usage.call_args.kwargs + kw = mock_usage_tracker.record_usage_nowait.call_args.kwargs assert kw["input_tokens"] == 11 assert kw["output_tokens"] == 4 assert kw["api_surface"] == "responses" @@ -104,7 +104,7 @@ def test_streaming_responses_records_usage_from_response_completed( assert b"response.completed" in out assert b"hi" in out - kw = mock_usage_tracker.record_usage.call_args.kwargs + kw = mock_usage_tracker.record_usage_nowait.call_args.kwargs assert kw["input_tokens"] == 12 assert kw["output_tokens"] == 3 assert kw["api_surface"] == "responses" @@ -172,7 +172,7 @@ def test_responses_upstream_error_returned_verbatim( ) assert r.status_code == 400 assert r.json()["error"]["message"] == "bad input" - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_streaming_responses_upstream_4xx_returns_json_not_sse( @@ -206,7 +206,7 @@ def test_streaming_responses_upstream_4xx_returns_json_not_sse( "application/json" ), f"expected JSON content-type, got {r.headers['content-type']}" assert r.json() == err - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_non_streaming_responses_web_search_uses_local_adapter_not_upstream( @@ -252,7 +252,7 @@ def test_non_streaming_responses_web_search_uses_local_adapter_not_upstream( assert not route.called assert mock_web_search_service.handle_request.called assert mock_response_context_store.save.called - kw = mock_usage_tracker.record_usage.call_args.kwargs + kw = mock_usage_tracker.record_usage_nowait.call_args.kwargs assert kw["api_surface"] == "responses" assert kw["input_tokens"] == 3 assert kw["output_tokens"] == 2 @@ -350,7 +350,7 @@ def test_streaming_responses_web_search_emits_local_responses_sse( assert "streamed answer" in out assert "event: response.completed" in out assert not route.called - kw = mock_usage_tracker.record_usage.call_args.kwargs + kw = mock_usage_tracker.record_usage_nowait.call_args.kwargs assert kw["api_surface"] == "responses" assert kw["input_tokens"] == 4 assert kw["output_tokens"] == 3 @@ -384,7 +384,7 @@ def test_streaming_responses_web_search_service_failure_returns_json_error( assert data["error"]["type"] == "api_error" assert "local failed" in data["error"]["message"] assert not route.called - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_non_streaming_responses_web_search_service_failure_returns_json_error( @@ -414,7 +414,7 @@ def test_non_streaming_responses_web_search_service_failure_returns_json_error( assert data["error"]["type"] == "api_error" assert "local failed" in data["error"]["message"] assert not route.called - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_non_streaming_responses_web_search_dependency_failure_returns_json_error( @@ -444,7 +444,7 @@ def test_non_streaming_responses_web_search_dependency_failure_returns_json_erro assert data["error"]["type"] == "api_error" assert "setup failed" in data["error"]["message"] assert not route.called - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_streaming_responses_web_search_dependency_failure_returns_json_error( @@ -475,7 +475,7 @@ def test_streaming_responses_web_search_dependency_failure_returns_json_error( assert data["error"]["type"] == "api_error" assert "setup failed" in data["error"]["message"] assert not route.called - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_responses_web_search_rejects_external_web_access_false( @@ -501,7 +501,7 @@ def test_responses_web_search_rejects_external_web_access_false( assert r.json()["error"]["type"] == "invalid_request_error" assert "external_web_access" in r.json()["error"]["message"] assert not route.called - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called @pytest.mark.parametrize("web_search_enabled", [False], indirect=True) @@ -532,4 +532,4 @@ def test_non_streaming_responses_web_search_disabled_skips_local_dependencies( assert "disabled" in data["error"]["message"].lower() assert not route.called assert not mock_web_search_service.handle_request.called - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called diff --git a/tests/integration/test_openai_passthrough/test_responses_crud.py b/tests/integration/test_openai_passthrough/test_responses_crud.py index 3692c3b..7fc99ee 100644 --- a/tests/integration/test_openai_passthrough/test_responses_crud.py +++ b/tests/integration/test_openai_passthrough/test_responses_crud.py @@ -10,7 +10,7 @@ def test_get_response_forwards_and_returns_body(client, respx_mock, mock_usage_t assert r.status_code == 200 assert r.json() == body # No usage logged for retrieval - assert not mock_usage_tracker.record_usage.called + assert not mock_usage_tracker.record_usage_nowait.called def test_delete_response_forwards(client, respx_mock): diff --git a/tests/unit/test_auth_cache.py b/tests/unit/test_auth_cache.py new file mode 100644 index 0000000..d3ad185 --- /dev/null +++ b/tests/unit/test_auth_cache.py @@ -0,0 +1,264 @@ +"""AuthMiddleware caches API key validation and keeps it off the event loop.""" +import asyncio +import threading +from unittest.mock import MagicMock + +import pytest +from starlette.requests import Request +from starlette.responses import PlainTextResponse + +VALID_INFO = {"api_key": "sk-valid", "user_id": "u1", "is_active": True} + + +def make_middleware(monkeypatch, cache_ttl_seconds=60): + from app.middleware import auth as auth_mod + + monkeypatch.setattr(auth_mod.settings, "require_api_key", True) + monkeypatch.setattr(auth_mod.settings, "master_api_key", "sk-master-secret") + + middleware = auth_mod.AuthMiddleware( + app=MagicMock(), + dynamodb_client=MagicMock(), + cache_ttl_seconds=cache_ttl_seconds, + ) + middleware.api_key_manager = MagicMock() + return middleware + + +def make_request(api_key=None, path="/v1/messages"): + from app.core.config import settings + + headers = [] + if api_key is not None: + headers.append((settings.api_key_header.lower().encode(), api_key.encode())) + scope = { + "type": "http", + "method": "POST", + "path": path, + "headers": headers, + "query_string": b"", + } + return Request(scope) + + +async def run_dispatch(middleware, api_key): + """Dispatch one request; return (response, api_key_info seen by handler).""" + seen = {} + + async def call_next(request): + seen["info"] = request.state.api_key_info + return PlainTextResponse("ok") + + response = await middleware.dispatch(make_request(api_key), call_next) + return response, seen.get("info") + + +def test_valid_key_hits_dynamo_once_within_ttl(monkeypatch): + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + middleware.api_key_manager.validate_api_key.return_value = VALID_INFO + + async def scenario(): + r1, info1 = await run_dispatch(middleware, "sk-valid") + r2, info2 = await run_dispatch(middleware, "sk-valid") + return r1, info1, r2, info2 + + r1, info1, r2, info2 = asyncio.run(scenario()) + + assert r1.status_code == 200 + assert r2.status_code == 200 + assert info1["user_id"] == "u1" + assert info2["user_id"] == "u1" + assert middleware.api_key_manager.validate_api_key.call_count == 1 + + +def test_cache_expiry_revalidates(monkeypatch): + from app.core import ttl_cache as ttl_mod + + now = 1000.0 + monkeypatch.setattr(ttl_mod.time, "monotonic", lambda: now) + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + middleware.api_key_manager.validate_api_key.return_value = VALID_INFO + + asyncio.run(run_dispatch(middleware, "sk-valid")) + monkeypatch.setattr(ttl_mod.time, "monotonic", lambda: now + 61) + asyncio.run(run_dispatch(middleware, "sk-valid")) + + assert middleware.api_key_manager.validate_api_key.call_count == 2 + + +def test_invalid_key_negative_cached_briefly(monkeypatch): + from app.core import ttl_cache as ttl_mod + + now = 1000.0 + monkeypatch.setattr(ttl_mod.time, "monotonic", lambda: now) + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + middleware.api_key_manager.validate_api_key.return_value = None + + r1, _ = asyncio.run(run_dispatch(middleware, "sk-bogus")) + r2, _ = asyncio.run(run_dispatch(middleware, "sk-bogus")) + assert r1.status_code == 401 + assert r2.status_code == 401 + # Second rejection served from the negative cache + assert middleware.api_key_manager.validate_api_key.call_count == 1 + + # Negative entries expire much sooner than positive ones (~5s, not 60s) + monkeypatch.setattr(ttl_mod.time, "monotonic", lambda: now + 6) + asyncio.run(run_dispatch(middleware, "sk-bogus")) + assert middleware.api_key_manager.validate_api_key.call_count == 2 + + +def test_ttl_zero_disables_cache(monkeypatch): + middleware = make_middleware(monkeypatch, cache_ttl_seconds=0) + middleware.api_key_manager.validate_api_key.return_value = VALID_INFO + + asyncio.run(run_dispatch(middleware, "sk-valid")) + asyncio.run(run_dispatch(middleware, "sk-valid")) + + assert middleware.api_key_manager.validate_api_key.call_count == 2 + + +def test_validation_runs_off_event_loop(monkeypatch): + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + validation_thread = {} + + def record_thread(api_key): + validation_thread["ident"] = threading.get_ident() + return VALID_INFO + + middleware.api_key_manager.validate_api_key.side_effect = record_thread + + async def scenario(): + loop_thread = threading.get_ident() + await run_dispatch(middleware, "sk-valid") + return loop_thread + + loop_thread = asyncio.run(scenario()) + + assert validation_thread["ident"] != loop_thread + + +def test_master_key_never_touches_dynamo(monkeypatch): + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + + response, info = asyncio.run(run_dispatch(middleware, "sk-master-secret")) + + assert response.status_code == 200 + assert info["is_master"] is True + middleware.api_key_manager.validate_api_key.assert_not_called() + + +def test_cached_info_not_shared_between_requests(monkeypatch): + """A handler mutating its api_key_info must not poison the cache.""" + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + middleware.api_key_manager.validate_api_key.return_value = dict(VALID_INFO) + + async def scenario(): + _, info1 = await run_dispatch(middleware, "sk-valid") + info1["user_id"] = "tampered" + _, info2 = await run_dispatch(middleware, "sk-valid") + return info2 + + info2 = asyncio.run(scenario()) + assert info2["user_id"] == "u1" + + +def test_validation_error_returns_401_but_is_not_cached(monkeypatch): + """A transient DynamoDB failure must not lock a good key out via cache.""" + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + middleware.api_key_manager.validate_api_key.side_effect = [ + RuntimeError("dynamo blip"), + VALID_INFO, + ] + + r1, _ = asyncio.run(run_dispatch(middleware, "sk-valid")) + r2, info2 = asyncio.run(run_dispatch(middleware, "sk-valid")) + + assert r1.status_code == 401 + assert r2.status_code == 200 + assert info2["user_id"] == "u1" + assert middleware.api_key_manager.validate_api_key.call_count == 2 + + +def test_manager_reraises_client_error_instead_of_returning_none(): + """DynamoDB throttling must surface as an error, not as "invalid key". + + If validate_api_key swallows ClientError and returns None, the + middleware negative-caches valid keys for 5s during any DynamoDB + throttling event — rolling 401 lockouts for legitimate traffic. + """ + from botocore.exceptions import ClientError + + from app.db.dynamodb import APIKeyManager + + manager = APIKeyManager.__new__(APIKeyManager) + manager.table = MagicMock() + manager.table.get_item.side_effect = ClientError( + {"Error": {"Code": "ProvisionedThroughputExceededException", "Message": "x"}}, + "GetItem", + ) + + with pytest.raises(ClientError): + manager.validate_api_key("sk-valid") + + +def test_nested_mutation_does_not_poison_cache(monkeypatch): + """DynamoDB items carry nested mutables (metadata); a handler mutating + one must not leak into other requests' cached copies.""" + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + middleware.api_key_manager.validate_api_key.return_value = { + **VALID_INFO, + "metadata": {"team": "alpha"}, + } + + async def scenario(): + _, info1 = await run_dispatch(middleware, "sk-valid") + info1["metadata"]["team"] = "tampered" + _, info2 = await run_dispatch(middleware, "sk-valid") + return info2 + + info2 = asyncio.run(scenario()) + assert info2["metadata"]["team"] == "alpha" + + +def test_validation_runs_on_dedicated_auth_executor(monkeypatch): + """Auth must not share the default executor with multi-second work + (Tavily web search, docker pulls) or cache misses queue behind them.""" + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + seen = {} + + def record_thread(api_key): + seen["name"] = threading.current_thread().name + return VALID_INFO + + middleware.api_key_manager.validate_api_key.side_effect = record_thread + asyncio.run(run_dispatch(middleware, "sk-valid")) + + assert seen["name"].startswith("auth-validate") + + +def test_concurrent_misses_share_one_lookup(monkeypatch): + """Single-flight: when a hot key's entry expires, concurrent misses + must coalesce into one DynamoDB read, not a thundering herd.""" + import time as time_mod + + middleware = make_middleware(monkeypatch, cache_ttl_seconds=60) + + def slow_validate(api_key): + time_mod.sleep(0.2) + return VALID_INFO + + middleware.api_key_manager.validate_api_key.side_effect = slow_validate + + async def scenario(): + results = await asyncio.gather( + run_dispatch(middleware, "sk-valid"), + run_dispatch(middleware, "sk-valid"), + run_dispatch(middleware, "sk-valid"), + ) + return results + + results = asyncio.run(scenario()) + + assert all(r.status_code == 200 for r, _ in results) + assert all(info["user_id"] == "u1" for _, info in results) + assert middleware.api_key_manager.validate_api_key.call_count == 1 diff --git a/tests/unit/test_model_mapping_cache.py b/tests/unit/test_model_mapping_cache.py new file mode 100644 index 0000000..c0468cd --- /dev/null +++ b/tests/unit/test_model_mapping_cache.py @@ -0,0 +1,138 @@ +"""ModelMappingManager caches mapping lookups across instances.""" +from unittest.mock import MagicMock + +import pytest +from botocore.exceptions import ClientError + +BEDROCK_ID = "global.anthropic.claude-sonnet-4-5-20250929-v1:0" +ANTHROPIC_ID = "claude-sonnet-4-5-20250929" + + +def make_manager(cache_ttl=300, get_item_response=None): + from app.db.dynamodb import ModelMappingManager + + m = ModelMappingManager.__new__(ModelMappingManager) + m.table = MagicMock() + if get_item_response is not None: + m.table.get_item.return_value = get_item_response + m._cache_ttl = cache_ttl + return m + + +@pytest.fixture(autouse=True) +def clear_shared_cache(): + from app.db.dynamodb import ModelMappingManager + + cache = getattr(ModelMappingManager, "_cache", None) + if cache is not None: + cache.clear() + yield + if cache is not None: + cache.clear() + + +def test_second_lookup_served_from_cache(): + m = make_manager(get_item_response={"Item": {"bedrock_model_id": BEDROCK_ID}}) + + assert m.get_mapping(ANTHROPIC_ID) == BEDROCK_ID + assert m.get_mapping(ANTHROPIC_ID) == BEDROCK_ID + assert m.table.get_item.call_count == 1 + + +def test_cache_shared_across_manager_instances(): + """Converters build a fresh manager per request; the cache must outlive them.""" + m1 = make_manager(get_item_response={"Item": {"bedrock_model_id": BEDROCK_ID}}) + m2 = make_manager() + + m1.get_mapping(ANTHROPIC_ID) + assert m2.get_mapping(ANTHROPIC_ID) == BEDROCK_ID + m2.table.get_item.assert_not_called() + + +def test_negative_result_cached(): + """Pass-through model IDs (no mapping) must not hit DynamoDB every request.""" + m = make_manager(get_item_response={}) + + assert m.get_mapping("arn:aws:bedrock:us-west-2::foundation-model/x") is None + assert m.get_mapping("arn:aws:bedrock:us-west-2::foundation-model/x") is None + assert m.table.get_item.call_count == 1 + + +def test_ttl_expiry_refetches(monkeypatch): + from app.core import ttl_cache as ttl_mod + + now = 1000.0 + monkeypatch.setattr(ttl_mod.time, "monotonic", lambda: now) + m = make_manager( + cache_ttl=300, get_item_response={"Item": {"bedrock_model_id": BEDROCK_ID}} + ) + + m.get_mapping(ANTHROPIC_ID) + monkeypatch.setattr(ttl_mod.time, "monotonic", lambda: now + 301) + m.get_mapping(ANTHROPIC_ID) + + assert m.table.get_item.call_count == 2 + + +def test_ttl_zero_disables_cache(): + m = make_manager( + cache_ttl=0, get_item_response={"Item": {"bedrock_model_id": BEDROCK_ID}} + ) + + m.get_mapping(ANTHROPIC_ID) + m.get_mapping(ANTHROPIC_ID) + + assert m.table.get_item.call_count == 2 + + +def test_client_error_not_cached(): + """A DynamoDB blip must not pin None for the whole TTL.""" + m = make_manager() + m.table.get_item.side_effect = [ + ClientError({"Error": {"Code": "Throttling", "Message": "x"}}, "GetItem"), + {"Item": {"bedrock_model_id": BEDROCK_ID}}, + ] + + assert m.get_mapping(ANTHROPIC_ID) is None + assert m.get_mapping(ANTHROPIC_ID) == BEDROCK_ID + assert m.table.get_item.call_count == 2 + + +def test_set_mapping_updates_cache_in_process(): + m = make_manager() + + m.set_mapping(ANTHROPIC_ID, BEDROCK_ID) + assert m.get_mapping(ANTHROPIC_ID) == BEDROCK_ID + m.table.get_item.assert_not_called() + + +def test_delete_mapping_invalidates_cache(): + m = make_manager(get_item_response={"Item": {"bedrock_model_id": BEDROCK_ID}}) + + m.get_mapping(ANTHROPIC_ID) # prime the cache + m.delete_mapping(ANTHROPIC_ID) + m.get_mapping(ANTHROPIC_ID) # must go back to DynamoDB + + assert m.table.get_item.call_count == 2 + + +def test_init_wires_ttl_from_settings(): + """Pin the real __init__ wiring (other tests bypass it via __new__).""" + from app.core.config import settings + from app.db.dynamodb import ModelMappingManager + + m = ModelMappingManager(MagicMock()) + + assert m._cache_ttl == settings.model_mapping_cache_ttl_seconds + + +def test_set_mapping_with_ttl_zero_does_not_cache(): + """With caching disabled, set_mapping must not plant a cache entry.""" + m = make_manager( + cache_ttl=0, get_item_response={"Item": {"bedrock_model_id": BEDROCK_ID}} + ) + + m.set_mapping(ANTHROPIC_ID, BEDROCK_ID) + m.get_mapping(ANTHROPIC_ID) + + assert m.table.get_item.call_count == 1 diff --git a/tests/unit/test_ttl_cache.py b/tests/unit/test_ttl_cache.py new file mode 100644 index 0000000..4a3c065 --- /dev/null +++ b/tests/unit/test_ttl_cache.py @@ -0,0 +1,136 @@ +"""Tests for the shared TTL cache utility.""" +import threading + +import pytest + + +@pytest.fixture +def cache(): + from app.core.ttl_cache import TTLCache + + return TTLCache(max_entries=4) + + +def test_miss_returns_no_hit(cache): + hit, value = cache.get("absent") + assert hit is False + assert value is None + + +def test_set_then_get_returns_value(cache): + cache.set("k", {"user": "u1"}, ttl_seconds=60) + hit, value = cache.get("k") + assert hit is True + assert value == {"user": "u1"} + + +def test_cached_none_is_a_hit(cache): + """Negative caching: a stored None must be distinguishable from a miss.""" + cache.set("invalid-key", None, ttl_seconds=5) + hit, value = cache.get("invalid-key") + assert hit is True + assert value is None + + +def test_expired_entry_is_a_miss(cache, monkeypatch): + from app.core import ttl_cache as mod + + now = 1000.0 + monkeypatch.setattr(mod.time, "monotonic", lambda: now) + cache.set("k", "v", ttl_seconds=60) + + monkeypatch.setattr(mod.time, "monotonic", lambda: now + 61) + hit, value = cache.get("k") + assert hit is False + assert value is None + + +def test_invalidate_removes_entry(cache): + cache.set("k", "v", ttl_seconds=60) + cache.invalidate("k") + hit, _ = cache.get("k") + assert hit is False + + +def test_invalidate_missing_key_is_noop(cache): + cache.invalidate("never-set") # must not raise + + +def test_eviction_purges_expired_before_clearing(cache, monkeypatch): + """When full, expired entries are purged so fresh ones survive.""" + from app.core import ttl_cache as mod + + now = 1000.0 + monkeypatch.setattr(mod.time, "monotonic", lambda: now) + cache.set("stale-1", "v", ttl_seconds=10) + cache.set("stale-2", "v", ttl_seconds=10) + cache.set("fresh", "v", ttl_seconds=600) + + # Advance past the stale TTLs, then fill to the max_entries=4 bound + monkeypatch.setattr(mod.time, "monotonic", lambda: now + 30) + cache.set("new-1", "v", ttl_seconds=600) + cache.set("new-2", "v", ttl_seconds=600) # triggers eviction of stale-* + + assert cache.get("fresh") == (True, "v") + assert cache.get("new-1") == (True, "v") + assert cache.get("new-2") == (True, "v") + + +def test_eviction_drops_soonest_expiring_when_nothing_expired(cache): + """Overflow must not wipe hot entries: short-TTL (negative-cache spam) + entries are dropped first, long-lived positive entries survive.""" + cache.set("short-1", "v", ttl_seconds=5) + cache.set("short-2", "v", ttl_seconds=5) + cache.set("long-1", "v", ttl_seconds=600) + cache.set("long-2", "v", ttl_seconds=600) + cache.set("overflow", "v", ttl_seconds=600) + + assert cache.get("overflow") == (True, "v") + assert cache.get("long-1") == (True, "v") + assert cache.get("long-2") == (True, "v") + # Bound is respected: never more entries than max_entries + assert len(cache) <= 4 + + +def test_exactly_at_expiry_is_miss(cache, monkeypatch): + """Pins the >= boundary: an entry is expired at exactly now + ttl.""" + from app.core import ttl_cache as mod + + now = 1000.0 + monkeypatch.setattr(mod.time, "monotonic", lambda: now) + cache.set("k", "v", ttl_seconds=60) + + monkeypatch.setattr(mod.time, "monotonic", lambda: now + 60) + assert cache.get("k") == (False, None) + + +def test_just_before_expiry_is_hit(cache, monkeypatch): + from app.core import ttl_cache as mod + + now = 1000.0 + monkeypatch.setattr(mod.time, "monotonic", lambda: now) + cache.set("k", "v", ttl_seconds=60) + + monkeypatch.setattr(mod.time, "monotonic", lambda: now + 59.999) + assert cache.get("k") == (True, "v") + + +def test_concurrent_access_does_not_corrupt(cache): + """Smoke test: concurrent set/get from threads must not raise.""" + errors = [] + + def worker(n): + try: + for i in range(200): + cache.set(f"k{n}-{i}", i, ttl_seconds=60) + cache.get(f"k{n}-{i}") + except Exception as e: # pragma: no cover + errors.append(e) + + threads = [threading.Thread(target=worker, args=(n,)) for n in range(4)] + for t in threads: + t.start() + for t in threads: + t.join() + + assert errors == [] diff --git a/tests/unit/test_usage_nowait.py b/tests/unit/test_usage_nowait.py new file mode 100644 index 0000000..f16c3a1 --- /dev/null +++ b/tests/unit/test_usage_nowait.py @@ -0,0 +1,134 @@ +"""UsageTracker.record_usage_nowait writes usage without blocking the caller.""" +import threading +from unittest.mock import MagicMock + +import pytest + + +@pytest.fixture +def tracker(): + from app.db.dynamodb import UsageTracker + + t = UsageTracker.__new__(UsageTracker) + t.table = MagicMock() + return t + + +USAGE_KWARGS = { + "api_key": "k", + "request_id": "r", + "model": "claude-sonnet-4-5-20250929", + "input_tokens": 10, + "output_tokens": 5, +} + + +def test_nowait_eventually_writes_same_item(tracker): + future = tracker.record_usage_nowait(**USAGE_KWARGS) + future.result(timeout=5) + + _, kwargs = tracker.table.put_item.call_args + item = kwargs["Item"] + assert item["api_key"] == "k" + assert item["request_id"] == "r" + assert item["input_tokens"] == 10 + assert item["output_tokens"] == 5 + + +def test_nowait_returns_before_write_happens(tracker): + """The caller must not wait for the DynamoDB write.""" + from app.db import dynamodb as mod + + release = threading.Event() + started = threading.Event() + + def blocker(): + started.set() + release.wait(timeout=5) + + # Occupy the single writer thread so the pending write cannot run yet + blocker_future = mod._get_usage_write_executor().submit(blocker) + try: + started.wait(timeout=5) + + future = tracker.record_usage_nowait(**USAGE_KWARGS) + assert tracker.table.put_item.call_count == 0 # caller returned, no write yet + finally: + release.set() + blocker_future.result(timeout=5) + future.result(timeout=5) + assert tracker.table.put_item.call_count == 1 + + +def test_nowait_swallows_write_errors(tracker): + """A failed usage write is telemetry loss, never a request failure.""" + tracker.table.put_item.side_effect = RuntimeError("dynamo down") + + future = tracker.record_usage_nowait(**USAGE_KWARGS) + future.result(timeout=5) # must not raise + + assert future.exception() is None + + +def test_nowait_rejects_unknown_kwargs_at_call_site(tracker): + """A misspelled kwarg must raise at the call site, not become a + TypeError swallowed inside the background thread (silent data loss).""" + with pytest.raises(TypeError): + tracker.record_usage_nowait(**USAGE_KWARGS, tokens_input=1) + + +def test_nowait_survives_shutdown_executor(tracker, monkeypatch): + """Call sites invoke nowait inside except/finally blocks; a submit-time + RuntimeError during shutdown must not mask the original exception.""" + from concurrent.futures import ThreadPoolExecutor + + from app.db import dynamodb as mod + + dead = ThreadPoolExecutor(max_workers=1) + dead.shutdown() + monkeypatch.setattr(mod, "_usage_write_executor", dead) + + future = tracker.record_usage_nowait(**USAGE_KWARGS) # must not raise + + assert future.done() + assert future.exception() is None + assert tracker.table.put_item.call_count == 0 + + +def test_nowait_drops_when_backlog_full(tracker, monkeypatch): + """Backpressure: when DynamoDB hangs and the queue backs up, new writes + are dropped (bounded memory) instead of accumulating without limit.""" + from app.db import dynamodb as mod + + monkeypatch.setattr(mod, "_MAX_PENDING_USAGE_WRITES", 0) + + future = tracker.record_usage_nowait(**USAGE_KWARGS) + + assert future.done() + assert future.exception() is None + assert tracker.table.put_item.call_count == 0 + + +def test_nowait_timestamp_captured_at_submit_with_ms_precision(tracker, monkeypatch): + """Timestamps must reflect when usage happened, not when the backlog + drained, and carry ms precision so same-second rows don't overwrite + each other (the usage table key is api_key + timestamp).""" + from app.db import dynamodb as mod + + monkeypatch.setattr(mod.time, "time", lambda: 1234.567) + future = tracker.record_usage_nowait(**USAGE_KWARGS) + future.result(timeout=5) + + _, kwargs = tracker.table.put_item.call_args + assert kwargs["Item"]["timestamp"] == "1234567" + + +def test_drain_usage_writes_flushes_pending(tracker): + """The shutdown hook must be able to flush queued writes.""" + from app.db import dynamodb as mod + + tracker.record_usage_nowait(**USAGE_KWARGS) + remaining = mod.drain_usage_writes(timeout=5) + + assert remaining == 0 + assert tracker.table.put_item.call_count == 1