Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions app/core/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:*).
Expand Down
44 changes: 43 additions & 1 deletion app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)


Expand Down
9 changes: 8 additions & 1 deletion app/middleware/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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()
Expand Down
8 changes: 7 additions & 1 deletion app/middleware/hpc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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("/")
Expand Down
6 changes: 5 additions & 1 deletion app/middleware/policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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)
Expand Down
95 changes: 95 additions & 0 deletions app/services/mcp_auth.py
Original file line number Diff line number Diff line change
@@ -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)
79 changes: 79 additions & 0 deletions app/services/mcp_server.py
Original file line number Diff line number Diff line change
@@ -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
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 3 additions & 0 deletions requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading
Loading