From 96410efedafc0cf015995913691dbbf2d4997353 Mon Sep 17 00:00:00 2001 From: man4ish Date: Sun, 4 Oct 2026 03:06:04 -0500 Subject: [PATCH 1/3] feat: hosted MCP server at /mcp (gap #10) Mounts the mcp SDK's Streamable HTTP transport at /mcp, gated by the same omni_sk_ API keys every other /v1 route already trusts (not the SDK's OAuth-flavored auth subsystem -- see mcp_auth.py's docstring for why publishing OAuth discovery metadata here would misrepresent what omnibioai-auth actually supports). Three tools (answer_with_citations, search_literature, list_domains) loop back to this gateway's own /v1/literature/* routes over HTTP, reusing 100% of existing rate-limit/quota/billing/BYOK logic instead of reimplementing it for a second transport. Fixed four bugs found while building this, each covered by a test: - /mcp mount was registered after the catch-all router and silently bypassed its own auth check - streamable_http_app()'s default host triggers DNS-rebinding protection that 404s every real (non-loopback) Host header (app/main.py) - the Streamable HTTP session manager's task group must be entered via session_manager.run() during the app's lifespan, not left to its own defaults (app/main.py) - FastAPI's root_path constructor override breaks Starlette's Mount child-root_path computation for everything nested under /mcp (app/services/mcp_auth.py) Still open: a full interactive OAuth 2.1 "connector" flow (discovery + dynamic client registration + PKCE) for generic MCP clients -- not buildable without fabricating capabilities omnibioai-auth doesn't have. Co-Authored-By: Claude Sonnet 5 --- app/core/config.py | 8 ++ app/main.py | 44 +++++++++- app/middleware/auth.py | 9 ++- app/middleware/hpc.py | 8 +- app/middleware/policy.py | 6 +- app/services/mcp_auth.py | 95 ++++++++++++++++++++++ app/services/mcp_server.py | 79 ++++++++++++++++++ requirements.txt | 3 + tests/test_mcp_auth.py | 121 ++++++++++++++++++++++++++++ tests/test_mcp_mount_integration.py | 98 ++++++++++++++++++++++ tests/test_mcp_server.py | 96 ++++++++++++++++++++++ 11 files changed, 563 insertions(+), 4 deletions(-) create mode 100644 app/services/mcp_auth.py create mode 100644 app/services/mcp_server.py create mode 100644 tests/test_mcp_auth.py create mode 100644 tests/test_mcp_mount_integration.py create mode 100644 tests/test_mcp_server.py diff --git a/app/core/config.py b/app/core/config.py index 093967b..8c79b69 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -40,6 +40,14 @@ class Config: # a silent fallback to the platform's own model. PROVIDER_KEY_REVEAL_SECRET = os.getenv("PROVIDER_KEY_REVEAL_SECRET", "") + # M17 (hosted MCP, design audit gap #10): the hosted /mcp endpoint's + # tool handlers call back into this gateway's own /v1/literature/* + # REST routes (see app/services/mcp_server.py) -- reusing their + # existing rate-limit/quota/idempotency/billing logic unchanged, + # rather than reimplementing it for a second transport. Loopback, + # not a path a browser or external client ever reaches directly. + SELF_BASE_URL = os.getenv("SELF_BASE_URL", "http://127.0.0.1:8080") + # Public /v1 API (app/routes/v1.py). Rate-limit, idempotency and quota # state lives under gateway:v1:* in the Redis at V1_REDIS_URL, which # needs its own ACL user (incr/expire/get/set/del/decr on gateway:v1:*). diff --git a/app/main.py b/app/main.py index e99f0d9..abbf571 100644 --- a/app/main.py +++ b/app/main.py @@ -17,11 +17,32 @@ from app.routes.auth_verify import router as auth_verify_router from app.routes.gateway import router from app.routes.v1 import router as v1_router +from app.services.mcp_auth import MCPBearerAuthASGIMiddleware +from app.services.mcp_server import build_mcp_server iam = IAMClient(Config.IAM_URL, Config.REDIS_URL) policy = PolicyClient(Config.POLICY_URL) hpc = HPCPolicyClient(Config.HPC_URL) +# M17 (hosted MCP, design audit gap #10): built once at import time, not +# inside lifespan -- app.mount() below needs the already-built +# Starlette app synchronously, and mcp_server.session_manager (used in +# lifespan further down) only exists after streamable_http_app() has +# been called on this same instance. +mcp_server = build_mcp_server() +mcp_streamable_http_app = mcp_server.streamable_http_app( + streamable_http_path="/", + # host is a label the SDK uses only to decide whether to + # *auto*-enable DNS-rebinding Host-header protection for a + # looks-like-a-local-demo app ("127.0.0.1"/"localhost"/"::1") -- + # not a bind address (uvicorn's own --host already controls that, + # see Dockerfile). Left at its own default, every real request's + # Host header (never literally "localhost") would silently 404 + # as if DNS rebinding protection were rejecting it, since the + # auto-enabled allowlist only accepts those three loopback names. + host="0.0.0.0", +) + _invalidation_task: asyncio.Task | None = None @@ -49,7 +70,15 @@ async def on_invalidate(user_id: str, token: str, api_key_hash: str = ""): async def lifespan(app: FastAPI): global _invalidation_task _invalidation_task = asyncio.create_task(_invalidation_loop()) - yield + # M17: the Streamable HTTP session manager's task group must be + # running before any /mcp request arrives, or every one fails with + # "Task group is not initialized" -- entering it here, open for + # this app's whole lifetime, is the SDK's own documented pattern + # for mounting an MCPServer into an existing ASGI app rather than + # running it standalone (see StreamableHTTPSessionManager.run's own + # docstring). + async with mcp_server.session_manager.run(): + yield if _invalidation_task: _invalidation_task.cancel() try: @@ -78,6 +107,19 @@ async def lifespan(app: FastAPI): # /v1 before the catch-all too, or /{service}/{path} would take "v1" as a # service name. app.include_router(v1_router) + +# M17 (hosted MCP, design audit gap #10): also before the catch-all, +# same reason -- /{service}/{path:path} would otherwise match +# service="mcp" and swallow every request here itself, never reaching +# this mount at all. Mounted, not routed through app.include_router, +# since Streamable HTTP's own ASGI app owns the full request/response +# (including a long-lived streaming session) for this path -- +# AuthMiddleware/PolicyMiddleware/HPCMiddleware above all skip "/mcp" +# explicitly (see each one's own _SKIP_PREFIXES) specifically so none +# of them buffers or blocks that stream; MCPBearerAuthASGIMiddleware +# here is this mount's own, equivalent gate (see app/services/mcp_auth.py). +app.mount("/mcp", MCPBearerAuthASGIMiddleware(mcp_streamable_http_app, iam)) + app.include_router(router) diff --git a/app/middleware/auth.py b/app/middleware/auth.py index 1299af5..0afb626 100644 --- a/app/middleware/auth.py +++ b/app/middleware/auth.py @@ -10,6 +10,13 @@ # security-audit, and toolserver all expose their own /docs pages. # Every actual API call still goes through the token check below. _SKIP_PATHS = {"/health", "/", "/version", "/docs", "/openapi.json"} +# M17: the mounted MCP Streamable HTTP app (app/services/mcp_server.py) +# is a prefix, not one exact path, and gates itself -- see +# app/services/mcp_auth.py's MCPBearerAuthASGIMiddleware, wrapping that +# mount directly. Skipped here (and in PolicyMiddleware/HPCMiddleware) +# so this request/response-oriented middleware never buffers or blocks +# that app's own long-lived streaming session. +_SKIP_PREFIXES = ("/mcp",) class AuthMiddleware(BaseHTTPMiddleware): @@ -18,7 +25,7 @@ def __init__(self, app, iam: IAMClient): self.iam = iam async def dispatch(self, request, call_next): - if request.url.path in _SKIP_PATHS: + if request.url.path in _SKIP_PATHS or request.url.path.startswith(_SKIP_PREFIXES): return await call_next(request) token = request.headers.get("Authorization", "").removeprefix("Bearer ").strip() diff --git a/app/middleware/hpc.py b/app/middleware/hpc.py index 0ea0cf4..0c479fa 100644 --- a/app/middleware/hpc.py +++ b/app/middleware/hpc.py @@ -5,6 +5,12 @@ from app.services.audit_client import build_audit_event, fire_audit _SKIP_PATHS = {"/health", "/", "/version"} +# M17: see app/middleware/auth.py's own _SKIP_PREFIXES comment. Already +# true in effect here too -- "mcp" is not a registered HPC compute +# service, so is_compute_service("mcp") already falls through to +# call_next below -- but explicit, like the other two middlewares, +# rather than relying on that incidentally being the case. +_SKIP_PREFIXES = ("/mcp",) class HPCMiddleware(BaseHTTPMiddleware): @@ -13,7 +19,7 @@ def __init__(self, app, hpc: HPCPolicyClient): self.hpc = hpc async def dispatch(self, request, call_next): - if request.url.path in _SKIP_PATHS: + if request.url.path in _SKIP_PATHS or request.url.path.startswith(_SKIP_PREFIXES): return await call_next(request) parts = request.url.path.strip("/").split("/") diff --git a/app/middleware/policy.py b/app/middleware/policy.py index 34d073c..981f74f 100644 --- a/app/middleware/policy.py +++ b/app/middleware/policy.py @@ -23,6 +23,10 @@ # an unmodeled synthetic path, not a real API-service route, so # exempting it doesn't touch actual API access control. _SKIP_PATHS = {"/health", "/", "/auth/verify", "/version", "/docs", "/openapi.json"} +# M17: see app/middleware/auth.py's own _SKIP_PREFIXES comment -- the +# mounted MCP app gates itself and must never be wrapped by this +# request/response-oriented middleware. +_SKIP_PREFIXES = ("/mcp",) class PolicyMiddleware(BaseHTTPMiddleware): @@ -31,7 +35,7 @@ def __init__(self, app, policy: PolicyClient): self.policy = policy async def dispatch(self, request, call_next): - if request.url.path in _SKIP_PATHS: + if request.url.path in _SKIP_PATHS or request.url.path.startswith(_SKIP_PREFIXES): return await call_next(request) user = getattr(request.state, "user", None) diff --git a/app/services/mcp_auth.py b/app/services/mcp_auth.py new file mode 100644 index 0000000..f0acdd8 --- /dev/null +++ b/app/services/mcp_auth.py @@ -0,0 +1,95 @@ +"""M17 (hosted MCP, design audit gap #10): Bearer-token gating for the +Streamable HTTP transport mounted at /mcp. + +Deliberately NOT the `mcp` SDK's own OAuth-flavored auth subsystem +(token_verifier + AuthSettings): that subsystem publishes OAuth +discovery metadata (.well-known/oauth-protected-resource) naming an +issuer, implying a standards-compliant authorization server a generic +MCP client could use to obtain a token via the interactive +authorization-code + PKCE "connector" flow. omnibioai-auth doesn't run +that flow (its own OAuth endpoints are client_credentials and one +first-party, platform-admin-only authorization_code path for LIMS, not +a general-purpose one) -- publishing discovery metadata claiming +otherwise would be misleading, not just incomplete. See this +milestone's checkpoint note for that still-open remainder of gap #10. + +What's actually implemented: an omni_sk_ API key, already a real, +independently-verifiable Bearer token (the exact mechanism every other +/v1 route already trusts), gates access to /mcp the same way. A raw +ASGI wrapper, not Starlette's BaseHTTPMiddleware, so a long-lived +Streamable HTTP session is never buffered -- just one header check +before the mounted MCP app ever runs. + +Also works around a real Starlette/FastAPI incompatibility found while +building this: app.main's app = FastAPI(..., root_path="/_svc/gateway") +makes FastAPI.__call__ force scope["root_path"] = "/_svc/gateway" on +every request. Starlette's Mount then computes this mount's own child +root_path as "/_svc/gateway" + "/mcp" (routing.Mount.matches: +"root_path": root_path + matched_path) -- but scope["path"] was never +actually prefixed with "/_svc/gateway" to begin with (root_path here +is FastAPI's own constructor override, not a prefix a real ASGI-aware +proxy stripped and reflected in both fields consistently), so +starlette.routing.get_route_path's path.startswith(root_path) check +fails and falls back to the full, unstripped path -- which then +matches nothing inside the mounted Streamable HTTP app. Every other +route in this service is a flat APIRouter-prefixed route, never a +nested Mount, so nothing had hit this before /mcp. Fixed below by +resetting root_path to just this mount's own prefix before forwarding, +which is what it would already correctly be if the outer app had no +root_path override at all. +""" +import contextvars + +from starlette.responses import JSONResponse + +from app.services.iam_client import is_api_key + +# Set once per inbound ASGI call by MCPBearerAuthASGIMiddleware, read by +# mcp_server.py's tool handlers to authenticate their own loopback call +# into this gateway's /v1/literature/* routes. contextvars are task- +# scoped, and each inbound HTTP call (including each one within a +# longer-lived Streamable HTTP session) runs in its own task, so this +# is never shared across two different callers' requests. +_current_api_key: contextvars.ContextVar[str | None] = contextvars.ContextVar( + "mcp_current_api_key", default=None, +) + + +def get_current_mcp_api_key() -> str | None: + return _current_api_key.get() + + +class MCPBearerAuthASGIMiddleware: + """Wraps the mounted MCP Streamable HTTP app. Rejects (401) any + request without a valid omni_sk_ key before it ever reaches the MCP + protocol handler; on success, stores the raw key for this request's + tool handlers to use via get_current_mcp_api_key().""" + + def __init__(self, app, iam): + self.app = app + self.iam = iam + + async def __call__(self, scope, receive, send): + if scope["type"] != "http": + return await self.app(scope, receive, send) + + # See this module's own docstring: undoes app.main's root_path + # override before it reaches Starlette's route matching inside + # the mounted app, where it would otherwise 404 every request. + scope = {**scope, "root_path": "/mcp"} + + headers = dict(scope.get("headers") or []) + auth_header = headers.get(b"authorization", b"").decode("latin-1") + token = auth_header.removeprefix("Bearer ").strip() + + if not token or not is_api_key(token): + response = JSONResponse({"error": "missing or invalid omni_sk_ API key"}, status_code=401) + return await response(scope, receive, send) + + user = await self.iam.validate_api_key(token) + if not user: + response = JSONResponse({"error": "invalid api key"}, status_code=401) + return await response(scope, receive, send) + + _current_api_key.set(token) + return await self.app(scope, receive, send) diff --git a/app/services/mcp_server.py b/app/services/mcp_server.py new file mode 100644 index 0000000..9ea012b --- /dev/null +++ b/app/services/mcp_server.py @@ -0,0 +1,79 @@ +"""M17 (hosted MCP, design audit gap #10): a hosted Streamable HTTP MCP +server at /mcp, so Claude, ChatGPT, and other MCP clients can reach the +public Literature AI API without each needing their own local stdio +process (omnibioai-sdk's omnibioai/mcp_server.py, which already exists +for that local case and is unaffected by this module). + +Every tool call here is a loopback HTTP call into this same gateway's +own, already-complete /v1/literature/* REST routes (app/routes/v1.py), +carrying the caller's own omni_sk_ key as its Authorization header -- +exactly the request shape a direct REST caller would send. This reuses +100% of the existing rate-limit, quota, idempotency, test-mode, BYOK, +and billing logic unchanged; nothing about metering or authorization is +reimplemented for this second transport. + +See app/services/mcp_auth.py's own module docstring for why this does +not use the `mcp` SDK's OAuth-flavored auth subsystem, and this +milestone's checkpoint note for what a full OAuth 2.1 "connector" flow +would still need beyond what's built here. +""" +from mcp.server.mcpserver import MCPServer + +from app.core.config import Config +from app.routes.gateway import proxy +from app.services.mcp_auth import get_current_mcp_api_key + +SERVER_NAME = "omnibioai" +INSTRUCTIONS = ( + "Biomedical literature tools backed by PubMed, billed through the caller's own " + "OmniBioAI organization. Use answer_with_citations for questions that need an " + "evidence-based answer; every claim carries its PubMed ID. Use search_literature " + "when only ranked source documents are needed, and list_domains to see which " + "research domains can be searched." +) + + +class MCPToolCallError(RuntimeError): + """Raised when the loopback REST call this tool wraps fails -- the + MCP framework surfaces this to the calling client as a tool error, + not a crashed session.""" + + +async def _call_self(path: str, method: str, body: dict | None = None) -> dict: + token = get_current_mcp_api_key() + if not token: + # Should be unreachable: MCPBearerAuthASGIMiddleware already + # rejected any request without one before this ever runs. A + # defensive backstop, not the primary check. + raise MCPToolCallError("No authenticated API key for this MCP session.") + status, response = await proxy.forward( + url=f"{Config.SELF_BASE_URL}{path}", + method=method, + headers={"Authorization": f"Bearer {token}"}, + body=body, + ) + if not 200 <= status < 300: + error = response.get("error") if isinstance(response, dict) else response + raise MCPToolCallError(f"OmniBioAI API error ({status}): {error}") + return response + + +def build_mcp_server() -> MCPServer: + server = MCPServer(SERVER_NAME, instructions=INSTRUCTIONS) + + @server.tool(description="Answer a biomedical question from PubMed abstracts, citing PMIDs. Billed per answer.") + async def answer_with_citations(question: str, domain: str = "default") -> dict: + return await _call_self("/v1/literature/answers", "POST", {"question": question, "domain": domain}) + + @server.tool( + description="Return ranked PubMed studies relevant to a question, with no generated answer -- " + "retrieval only. Billed per call, at a lower rate than an answer." + ) + async def search_literature(question: str, domain: str = "default") -> dict: + return await _call_self("/v1/literature/search", "POST", {"question": question, "domain": domain}) + + @server.tool(description="List the research domains that can be searched. Free.") + async def list_domains() -> dict: + return await _call_self("/v1/literature/domains", "GET") + + return server diff --git a/requirements.txt b/requirements.txt index 98b10ee..e284834 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,6 +3,9 @@ uvicorn httpx redis[asyncio] pyjwt +# M17 (hosted MCP, design audit gap #10): the Streamable HTTP transport +# mounted at /mcp (app/services/mcp_server.py). +mcp>=2.3.0,<3.0 # Keep the IAM client's transitive cryptography dependency on the validated # ARM64-compatible range; newer wheels use unsupported instructions here. cryptography==43.0.3 diff --git a/tests/test_mcp_auth.py b/tests/test_mcp_auth.py new file mode 100644 index 0000000..327ba4d --- /dev/null +++ b/tests/test_mcp_auth.py @@ -0,0 +1,121 @@ +"""app/services/mcp_auth.py: Bearer-token gating for the mounted /mcp +Streamable HTTP app. Raw ASGI-level tests -- MCPBearerAuthASGIMiddleware +is deliberately not Starlette's BaseHTTPMiddleware (see its own +docstring), so it's exercised at the scope/receive/send level directly +rather than through a TestClient wrapping a real downstream app. +""" +from unittest.mock import AsyncMock + +import pytest + +from app.services.mcp_auth import MCPBearerAuthASGIMiddleware, get_current_mcp_api_key + +KEY = "omni_sk_" + "a" * 40 + + +def _http_scope(auth_header: str | None): + headers = [] + if auth_header is not None: + headers.append((b"authorization", auth_header.encode())) + return {"type": "http", "method": "GET", "path": "/mcp", "headers": headers} + + +class _RecordingReceive: + async def __call__(self): + return {"type": "http.disconnect"} + + +class _RecordingSend: + def __init__(self): + self.messages = [] + + async def __call__(self, message): + self.messages.append(message) + + +def _status_of(send: _RecordingSend) -> int: + start = next(m for m in send.messages if m["type"] == "http.response.start") + return start["status"] + + +@pytest.fixture +def inner_app(): + """Records that it was called, and what get_current_mcp_api_key() + returned at the moment it ran -- proving the contextvar was set + *before* the inner app runs, not just that auth passed.""" + calls = [] + + async def app(scope, receive, send): + calls.append(get_current_mcp_api_key()) + await send({"type": "http.response.start", "status": 200, "headers": []}) + await send({"type": "http.response.body", "body": b"{}"}) + + app.calls = calls + return app + + +@pytest.fixture +def iam(): + return AsyncMock() + + +def test_missing_authorization_header_is_rejected_without_calling_inner_app(inner_app, iam): + middleware = MCPBearerAuthASGIMiddleware(inner_app, iam) + send = _RecordingSend() + import asyncio + asyncio.run(middleware(_http_scope(None), _RecordingReceive(), send)) + + assert _status_of(send) == 401 + assert inner_app.calls == [] + iam.validate_api_key.assert_not_called() + + +def test_non_api_key_bearer_token_is_rejected(inner_app, iam): + middleware = MCPBearerAuthASGIMiddleware(inner_app, iam) + send = _RecordingSend() + import asyncio + asyncio.run(middleware(_http_scope("Bearer not-an-api-key"), _RecordingReceive(), send)) + + assert _status_of(send) == 401 + assert inner_app.calls == [] + iam.validate_api_key.assert_not_called() + + +def test_invalid_api_key_is_rejected(inner_app, iam): + iam.validate_api_key.return_value = None + middleware = MCPBearerAuthASGIMiddleware(inner_app, iam) + send = _RecordingSend() + import asyncio + asyncio.run(middleware(_http_scope(f"Bearer {KEY}"), _RecordingReceive(), send)) + + assert _status_of(send) == 401 + assert inner_app.calls == [] + + +def test_valid_api_key_reaches_the_inner_app_with_the_key_set(inner_app, iam): + iam.validate_api_key.return_value = {"user_id": "5", "org_id": "42"} + middleware = MCPBearerAuthASGIMiddleware(inner_app, iam) + send = _RecordingSend() + import asyncio + asyncio.run(middleware(_http_scope(f"Bearer {KEY}"), _RecordingReceive(), send)) + + assert _status_of(send) == 200 + assert inner_app.calls == [KEY] + iam.validate_api_key.assert_awaited_once_with(KEY) + + +def test_contextvar_is_unset_outside_any_request(): + assert get_current_mcp_api_key() is None + + +def test_non_http_scope_passes_through_untouched(iam): + calls = [] + + async def app(scope, receive, send): + calls.append(scope["type"]) + + middleware = MCPBearerAuthASGIMiddleware(app, iam) + import asyncio + asyncio.run(middleware({"type": "lifespan"}, _RecordingReceive(), _RecordingSend())) + assert calls == ["lifespan"] + iam.validate_api_key.assert_not_called() diff --git a/tests/test_mcp_mount_integration.py b/tests/test_mcp_mount_integration.py new file mode 100644 index 0000000..8b7f2e1 --- /dev/null +++ b/tests/test_mcp_mount_integration.py @@ -0,0 +1,98 @@ +"""Integration-level confirmation that /mcp is reachable through the +real app (app.main), bypasses Auth/Policy/HPC middleware exactly as +intended, is gated by MCPBearerAuthASGIMiddleware's own check instead, +and that a real MCP protocol handshake actually completes. + +That last point matters: this app sets +FastAPI(root_path="/_svc/gateway") for Studio's reverse-proxy URL +generation, and FastAPI.__call__ forces scope["root_path"] to that +value on every request. Starlette's Mount then computes this mount's +own child root_path as "/_svc/gateway" + "/mcp", which +starlette.routing.get_route_path can no longer correctly subtract from +the (never actually so-prefixed) scope["path"] -- every request inside +the mount 404s unless that's corrected first (see +app/services/mcp_auth.py's own docstring for the full mechanism and +the fix). A test that only checked "not 401" would have silently passed +throughout -- it genuinely did, during development -- while the mount +was completely broken underneath; the handshake test below exists +specifically so that regressing this again fails loudly. +""" +import asyncio +from unittest.mock import AsyncMock, patch + +import httpx + +import app.main as _main_mod + +KEY = "omni_sk_" + "a" * 40 + +INITIALIZE_BODY = { + "jsonrpc": "2.0", + "method": "initialize", + "id": 1, + "params": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "clientInfo": {"name": "test-client", "version": "1.0"}, + }, +} +MCP_HEADERS = {"Accept": "application/json, text/event-stream"} + + +def _mcp_post(client, body, headers=None): + return client.request("POST", "/mcp/", json=body, headers={**MCP_HEADERS, **(headers or {})}, timeout=5) + + +def test_mcp_without_a_token_returns_this_mounts_own_401_not_authmiddlewares(client): + """AuthMiddleware's own 401 body is {"error": "missing token"} -- + getting MCPBearerAuthASGIMiddleware's distinct message instead + proves AuthMiddleware's _SKIP_PREFIXES actually skipped this path, + rather than it happening to also reject with the same status.""" + resp = _mcp_post(client, INITIALIZE_BODY) + assert resp.status_code == 401 + assert resp.json() == {"error": "missing or invalid omni_sk_ API key"} + + +def test_mcp_with_a_non_api_key_bearer_is_rejected(client): + resp = _mcp_post(client, INITIALIZE_BODY, headers={"Authorization": "Bearer not-an-api-key"}) + assert resp.status_code == 401 + + +def test_mcp_with_an_invalid_api_key_is_rejected(client): + with patch.object(_main_mod.iam, "validate_api_key", AsyncMock(return_value=None)): + resp = _mcp_post(client, INITIALIZE_BODY, headers={"Authorization": f"Bearer {KEY}"}) + assert resp.status_code == 401 + + +def test_mcp_initialize_handshake_succeeds_with_a_valid_api_key(): + """The real protocol-level proof: a valid key completes a genuine + MCP initialize exchange and gets back this server's own name and + capabilities -- not just "some non-401 response." + + Bypasses the module's shared `client` fixture: that fixture's + SyncASGIClient calls asyncio.run() per request without ever driving + the app's own ASGI lifespan, so mcp_server.session_manager's task + group (entered inside app.main.lifespan, see that module) is never + initialized -- any real Streamable HTTP request fails with "Task + group is not initialized." Entering session_manager.run() and + issuing the request inside the same asyncio.run() call is what + app.main.lifespan would already be doing for us in production. + """ + + async def _do_request(): + async with _main_mod.mcp_server.session_manager.run(): + transport = httpx.ASGITransport(app=_main_mod.app) + async with httpx.AsyncClient(transport=transport, base_url="http://testserver") as c: + return await c.post( + "/mcp/", + json=INITIALIZE_BODY, + headers={**MCP_HEADERS, "Authorization": f"Bearer {KEY}"}, + ) + + with patch.object(_main_mod.iam, "validate_api_key", AsyncMock(return_value={"user_id": "5", "org_id": "42"})): + resp = asyncio.run(_do_request()) + + assert resp.status_code == 200 + body = resp.text + assert '"serverInfo"' in body + assert '"omnibioai"' in body diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py new file mode 100644 index 0000000..91e75ff --- /dev/null +++ b/tests/test_mcp_server.py @@ -0,0 +1,96 @@ +"""app/services/mcp_server.py: the hosted MCP tools, each a loopback +call into this gateway's own /v1/literature/* REST routes carrying the +caller's own omni_sk_ key -- reusing that existing rate-limit/quota/ +billing logic unchanged, not reimplementing it for this transport. +""" +import asyncio +from unittest.mock import AsyncMock, patch + +import pytest + +from app.services.mcp_auth import _current_api_key +from app.services.mcp_server import MCPToolCallError, _call_self, build_mcp_server + +KEY = "omni_sk_" + "a" * 40 + + +@pytest.fixture(autouse=True) +def api_key_context(): + """Simulates MCPBearerAuthASGIMiddleware having already set the + contextvar for this "request" -- these tests exercise the tool + logic in isolation, not the middleware that sets it up.""" + token = _current_api_key.set(KEY) + yield + _current_api_key.reset(token) + + +def test_build_mcp_server_registers_the_three_tools(): + server = build_mcp_server() + tools = asyncio.run(server.list_tools()) + names = {t.name for t in tools} + assert names == {"answer_with_citations", "search_literature", "list_domains"} + + +class TestCallSelf: + def test_forwards_the_api_key_as_a_bearer_token(self): + forward = AsyncMock(return_value=(200, {"answer": "ok"})) + with patch("app.services.mcp_server.proxy.forward", forward): + result = asyncio.run(_call_self("/v1/literature/answers", "POST", {"question": "q"})) + assert result == {"answer": "ok"} + kwargs = forward.call_args.kwargs + assert kwargs["headers"] == {"Authorization": f"Bearer {KEY}"} + assert kwargs["method"] == "POST" + assert kwargs["body"] == {"question": "q"} + assert kwargs["url"].endswith("/v1/literature/answers") + + def test_raises_mcp_tool_call_error_on_a_non_2xx_response(self): + forward = AsyncMock(return_value=(402, {"error": {"type": "quota_exceeded", "message": "no quota"}})) + with patch("app.services.mcp_server.proxy.forward", forward): + with pytest.raises(MCPToolCallError): + asyncio.run(_call_self("/v1/literature/answers", "POST", {"question": "q"})) + + def test_raises_without_an_authenticated_key(self): + """Defensive backstop -- MCPBearerAuthASGIMiddleware should + already have rejected this request, but the tool layer must + never silently call out with no Authorization header at all.""" + reset_token = _current_api_key.set(None) + try: + with pytest.raises(MCPToolCallError): + asyncio.run(_call_self("/v1/literature/answers", "POST", {"question": "q"})) + finally: + _current_api_key.reset(reset_token) + + +class TestToolHandlers: + def _tool(self, server, name): + tools = {t.name: t for t in asyncio.run(server.list_tools())} + assert name in tools + return server + + def test_answer_with_citations_calls_the_answers_endpoint(self): + server = build_mcp_server() + forward = AsyncMock(return_value=(200, {"answer": "TP53 is a tumor suppressor."})) + with patch("app.services.mcp_server.proxy.forward", forward): + result = asyncio.run(server.call_tool("answer_with_citations", {"question": "What does TP53 do?"})) + kwargs = forward.call_args.kwargs + assert kwargs["url"].endswith("/v1/literature/answers") + assert kwargs["body"] == {"question": "What does TP53 do?", "domain": "default"} + + def test_search_literature_calls_the_search_endpoint(self): + server = build_mcp_server() + forward = AsyncMock(return_value=(200, {"results": []})) + with patch("app.services.mcp_server.proxy.forward", forward): + asyncio.run(server.call_tool("search_literature", {"question": "q", "domain": "Oncology"})) + kwargs = forward.call_args.kwargs + assert kwargs["url"].endswith("/v1/literature/search") + assert kwargs["body"] == {"question": "q", "domain": "Oncology"} + + def test_list_domains_calls_the_domains_endpoint_with_no_body(self): + server = build_mcp_server() + forward = AsyncMock(return_value=(200, {"domains": []})) + with patch("app.services.mcp_server.proxy.forward", forward): + asyncio.run(server.call_tool("list_domains", {})) + kwargs = forward.call_args.kwargs + assert kwargs["url"].endswith("/v1/literature/domains") + assert kwargs["method"] == "GET" + assert kwargs["body"] is None From 1f72d89d8358c0a73e1159dfa34feedc5a5be20e Mon Sep 17 00:00:00 2001 From: man4ish Date: Sun, 4 Oct 2026 03:11:27 -0500 Subject: [PATCH 2/3] fix: remove unused variable flagged by ruff in test_mcp_server.py Co-Authored-By: Claude Sonnet 5 --- tests/test_mcp_server.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py index 91e75ff..af27faf 100644 --- a/tests/test_mcp_server.py +++ b/tests/test_mcp_server.py @@ -71,7 +71,7 @@ def test_answer_with_citations_calls_the_answers_endpoint(self): server = build_mcp_server() forward = AsyncMock(return_value=(200, {"answer": "TP53 is a tumor suppressor."})) with patch("app.services.mcp_server.proxy.forward", forward): - result = asyncio.run(server.call_tool("answer_with_citations", {"question": "What does TP53 do?"})) + asyncio.run(server.call_tool("answer_with_citations", {"question": "What does TP53 do?"})) kwargs = forward.call_args.kwargs assert kwargs["url"].endswith("/v1/literature/answers") assert kwargs["body"] == {"question": "What does TP53 do?", "domain": "default"} From d784990ac249e09532758a8c91afdfd9f2328ee2 Mon Sep 17 00:00:00 2001 From: man4ish Date: Sun, 4 Oct 2026 03:12:47 -0500 Subject: [PATCH 3/3] fix: declare mcp dependency in pyproject.toml, not just requirements.txt CI installs via `pip install -e ".[dev]"` from pyproject.toml -- requirements.txt isn't read by CI or the Dockerfile (both use pyproject.toml), so the earlier requirements.txt-only addition left CI's environment without the mcp package, failing every test that imports app.main. Co-Authored-By: Claude Sonnet 5 --- pyproject.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/pyproject.toml b/pyproject.toml index a48f155..c05f8bb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -13,6 +13,7 @@ dependencies = [ "redis[asyncio]", "pyjwt", "cryptography==43.0.3", + "mcp>=2.3.0,<3.0", # IAM Foundation gateway integration: app/services/iam_client.py uses # AsyncIAMClient for RS256/JWKS/HS256 token verification. Pinned direct