Add AI engine load balancer authentication
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user