Add AI engine load balancer authentication

This commit is contained in:
Anthony Stirling
2026-05-27 12:17:15 +01:00
parent 4564ed5bec
commit 00a36ba0f5
15 changed files with 684 additions and 242 deletions
+19 -9
View File
@@ -1,9 +1,9 @@
from __future__ import annotations
import logging
from contextlib import asynccontextmanager
from typing import Annotated
from fastapi import Depends, FastAPI
from fastapi import FastAPI
from pydantic_ai import Agent
from pydantic_ai.models.instrumented import InstrumentationSettings
@@ -16,7 +16,7 @@ from stirling.agents import (
)
from stirling.agents.ledger import MathAuditorAgent
from stirling.agents.pdf_comment import PdfCommentAgent
from stirling.api.middleware import UserIdMiddleware
from stirling.api.middleware import EngineAuthMiddleware, UserIdMiddleware
from stirling.api.routes import (
agent_draft_router,
document_router,
@@ -63,7 +63,21 @@ async def lifespan(fast_api: FastAPI):
app = FastAPI(title="Stirling AI Engine", lifespan=lifespan, version="0.1.0")
try:
_engine_shared_secret = load_settings().engine_shared_secret or ""
except (AttributeError, KeyError) as cfg_err:
raise RuntimeError(
"engine_shared_secret missing from settings; ensure STIRLING_ENGINE_SHARED_SECRET "
"is declared in the env (blank value is allowed for dev mode)."
) from cfg_err
if not _engine_shared_secret:
logging.getLogger(__name__).warning(
"STIRLING_ENGINE_SHARED_SECRET is blank - running in dev (open) mode."
)
app.add_middleware(UserIdMiddleware)
app.add_middleware(EngineAuthMiddleware, expected_secret=_engine_shared_secret)
app.include_router(orchestrator_router)
app.include_router(pdf_edit_router)
app.include_router(pdf_question_router)
@@ -75,9 +89,5 @@ app.include_router(pdf_comments_router)
@app.get("/health", response_model=HealthResponse)
async def healthcheck(settings: Annotated[AppSettings, Depends(load_settings)]) -> HealthResponse:
return HealthResponse(
status="ok",
smart_model=settings.smart_model_name,
fast_model=settings.fast_model_name,
)
async def healthcheck() -> HealthResponse:
return HealthResponse(status="ok")
+24 -2
View File
@@ -1,16 +1,20 @@
from __future__ import annotations
import hmac
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
from starlette.requests import Request
from starlette.responses import Response
from starlette.responses import JSONResponse, Response
from stirling.services.tracking import current_user_id
_USER_ID_HEADER = "X-User-Id"
_ENGINE_AUTH_HEADER = "X-Engine-Auth"
_HEALTH_PATHS = {"/health", "/healthz", "/readyz"}
class UserIdMiddleware(BaseHTTPMiddleware):
"""Extract X-User-Id header and set it as the current user for PostHog tracking."""
"""Set X-User-Id (stamped by the trusted Java proxy) as request context."""
async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
user_id = request.headers.get(_USER_ID_HEADER)
@@ -21,3 +25,21 @@ class UserIdMiddleware(BaseHTTPMiddleware):
finally:
current_user_id.reset(token)
return await call_next(request)
class EngineAuthMiddleware(BaseHTTPMiddleware):
"""Validate shared-secret header. Blank secret = dev mode (open); health probes exempt."""
def __init__(self, app, expected_secret: str) -> None:
super().__init__(app)
self._expected_secret = expected_secret or ""
# Precompute the bytes form once so per-request work is just the constant-time compare.
self._expected_secret_bytes = self._expected_secret.encode("utf-8")
async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
if request.url.path in _HEALTH_PATHS or not self._expected_secret:
return await call_next(request)
presented = request.headers.get(_ENGINE_AUTH_HEADER, "")
if not hmac.compare_digest(presented.encode("utf-8"), self._expected_secret_bytes):
return JSONResponse({"detail": "engine authentication failed"}, status_code=401)
return await call_next(request)
+3
View File
@@ -98,6 +98,9 @@ class AppSettings(BaseSettings):
posthog_api_key: str = Field(validation_alias="STIRLING_POSTHOG_API_KEY")
posthog_host: str = Field(validation_alias="STIRLING_POSTHOG_HOST")
# Engine shared-secret for cluster deployments. Blank = dev mode (open).
engine_shared_secret: str = Field(default="", validation_alias="STIRLING_ENGINE_SHARED_SECRET")
def _configure_logging(level_name: str, log_file: str, http_debug: bool) -> None:
"""Configure the ``stirling`` logger hierarchy."""
-2
View File
@@ -5,5 +5,3 @@ from stirling.models import ApiModel
class HealthResponse(ApiModel):
status: str
smart_model: str
fast_model: str
@@ -0,0 +1,59 @@
"""Tests for EngineAuthMiddleware - shared-secret gating and health probe exemption."""
from __future__ import annotations
from fastapi import FastAPI
from fastapi.testclient import TestClient
from stirling.api.middleware import EngineAuthMiddleware
def _app(expected_secret: str) -> FastAPI:
app = FastAPI()
app.add_middleware(EngineAuthMiddleware, expected_secret=expected_secret)
@app.get("/health")
async def health():
return {"status": "ok"}
@app.get("/v1/agent")
async def agent():
return {"ok": True}
return app
def test_correct_secret_allows_request():
client = TestClient(_app("s3cret"))
res = client.get("/v1/agent", headers={"X-Engine-Auth": "s3cret"})
assert res.status_code == 200
assert res.json() == {"ok": True}
def test_missing_header_rejected_with_401():
client = TestClient(_app("s3cret"))
res = client.get("/v1/agent")
assert res.status_code == 401
def test_wrong_header_rejected_with_401():
client = TestClient(_app("s3cret"))
res = client.get("/v1/agent", headers={"X-Engine-Auth": "wrong"})
assert res.status_code == 401
def test_blank_secret_dev_mode_allows_unauthenticated():
client = TestClient(_app(""))
res = client.get("/v1/agent")
assert res.status_code == 200
def test_health_endpoint_exempt_from_auth():
client = TestClient(_app("s3cret"))
res = client.get("/health")
assert res.status_code == 200
def test_health_endpoint_exempt_in_dev_mode():
client = TestClient(_app(""))
res = client.get("/health")
assert res.status_code == 200
+50
View File
@@ -0,0 +1,50 @@
"""Tests for UserIdMiddleware - X-User-Id propagates as context.
X-Tenant-Id is deliberately NOT accepted from the wire: the Java proxy does not stamp
one and a client-supplied tenant id would be spoofable. If/when tenant scoping arrives
it must be derived from the server-side security context, not from request headers.
"""
from __future__ import annotations
from fastapi import FastAPI, Request
from fastapi.testclient import TestClient
from stirling.api.middleware import UserIdMiddleware
from stirling.services.tracking import current_user_id
def _app() -> FastAPI:
app = FastAPI()
app.add_middleware(UserIdMiddleware)
@app.get("/me")
async def me(request: Request):
return {
"user_id": current_user_id.get(),
"tenant_id": getattr(request.state, "tenant_id", None),
}
return app
def test_user_id_header_is_propagated():
client = TestClient(_app())
res = client.get("/me", headers={"X-User-Id": "alice"})
assert res.status_code == 200
assert res.json()["user_id"] == "alice"
def test_tenant_id_header_is_ignored():
client = TestClient(_app())
res = client.get("/me", headers={"X-User-Id": "alice", "X-Tenant-Id": "acme"})
assert res.status_code == 200
body = res.json()
assert body["user_id"] == "alice"
assert body["tenant_id"] is None
def test_missing_user_id_returns_no_user_in_context():
client = TestClient(_app())
res = client.get("/me")
assert res.status_code == 200
assert res.json()["user_id"] in ("", None)