mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
feat(proxy): resolve root_path per request from a configured prefix list (SERVER_ROOT_PATHS) (#35935)
* feat(proxy): resolve root_path per request from SERVER_ROOT_PATHS One deployment can encode exactly one client-visible URL path prefix today: SERVER_ROOT_PATH is a scalar stamped onto the app at startup, so a pod fronting several ingress prefixes 404s every prefix but one before any handler runs, and MCP OAuth discovery can emit only one prefix's URLs (RFC 9728 section 3 exact-match fails for the rest). Add an opt-in outermost ASGI middleware that matches the request path against a configured prefix list (SERVER_ROOT_PATHS, comma-separated) on a segment boundary and sets scope["root_path"] for that request only. Everything downstream is stock Starlette: route matching strips root_path so routes stay registered root-relative, and request.base_url re-includes it, so the discovery documents' resource and the 401 challenges' resource_metadata land under the prefix the client actually called — with no discovery-builder changes. LazyFeatureMiddleware now strips the scope root_path (falling back to the cached SERVER_ROOT_PATH scalar) before feature prefix matching, so lazily-registered routers — the MCP OAuth discovery router among them — load under per-request prefixes. Follow-up to the routing discussion on #35226; composes with, but does not depend on, #35576. * fix(proxy): import Sequence from collections.abc (ruff UP035 strict-budget gate) * review(greptile): trim implementation commentary; fixture-own MCP registry state in tests Addresses both P2s from the first Greptile pass: - per_request_root_path_middleware.py (and the related _lazy_features / proxy_server comments) cut down to the constraints the code cannot express, per repo comment guidance - the new discovery tests no longer clear/repopulate the shared MCP registry inline; a fixture snapshots it, hands the test an empty registry, and restores it afterwards so no state leaks between cases * fix(lint): mutable-ok marker on the prefix accumulator (LIT002 type-discipline gate) * fix(proxy): tie 401 challenges and get_custom_url to the per-request root_path The per-request root_path middleware sets scope["root_path"] to the prefix the client actually called, but the OAuth 401 challenges (raise_user_oauth_challenge / raise_token_exchange_challenge) still built their resource_metadata from SERVER_ROOT_PATH. On a pod fronting several prefixes, the challenge advertised a discovery URL under a different prefix than the discovery document served — the two disagreed on where the resource metadata lives, and a strict RFC 9728 client refused the challenge. Route the challenges through a small ContextVar the middleware populates so they read the same effective root_path Starlette resolves the request under. The same accessor fixes get_custom_url: when a request lives under a SERVER_ROOT_PATHS-matched prefix, request.base_url already carries it, so appending the SERVER_ROOT_PATH scalar on top produced e.g. /tenant-a/legacy/sso/callback — a path that does not exist. Reading the per-request prefix instead (and relying on join_paths's tail-dedup) keeps SSO login/callback URLs under one prefix — the one the request actually arrived on. Fallback: outside a request (module-load-time UI URL builders, background tasks) the ContextVar is unset and the accessor reads SERVER_ROOT_PATH, matching get_server_root_path() so scalar-only deployments are byte-identical. * fix(mcp): challenge URL under per-request prefix must route, and mock parity Two follow-ups to the review fix that made the 401 challenge use the per-request root_path: 1. oauth_protected_resource_path must pick the URL structure that actually routes for the mechanism in use: - The scalar SERVER_ROOT_PATH deployment registers the well-known routes with the prefix INSERTED (via well_known_root_suffix at import time), matching RFC 8414 §3. The challenge URL must use the same insertion or a client fetching it 404s. - The per-request SERVER_ROOT_PATHS deployment can't register routes per prefix; PerRequestRootPathMiddleware strips the prefix from scope["path"] and the router matches the un-inserted route. The URL must place the prefix BEFORE .well-known so the strip leaves a matching path. The previous fix used the insertion form for both, which 404'd the discovery fetch on the per-request path — the discovery doc and the challenge would then disagree on where the resource metadata lives, the very failure the review flagged. End-to-end verified: the URL the challenge advertises routes and the doc's `resource` field equals the URL the client originally called (RFC 9728 §3). 2. get_request_root_path now delegates its fallback through get_server_root_path() instead of reading the env directly, so every existing `monkeypatch.setattr("litellm.proxy.utils.get_server_root_path"` test override keeps working. This unstubbed the mock on the /v2/login test that failed on the last CI run. Plus the lint budget: annotate the local accumulator Final, tag the scope["root_path"] rewrite as an intentional ASGI-contract mutation, tag the reused `path`/`root_path` rebinds in LazyFeatureMiddleware, and add reason strings to the two new PLC0415 lazy-import noqas. * test(mcp): pin the reviewer's expected end-state — challenge URL routes, resource matches called URL End-to-end regression test that mounts the discoverable router + the per-request root_path middleware, hits an MCP endpoint that raises raise_user_oauth_challenge, fetches the resource_metadata URL the challenge advertises, and checks the returned document's `resource` equals the URL the client originally called (RFC 9728 §3 exact match). Covers /tenant-a, /tenant-b, and the unprefixed path on the same app so a regression on any prefix — challenge URL 404s, or doc emits a different prefix than the client called — fails at this test rather than in a strict MCP client's discovery. --------- Co-authored-by: gym-cmd <186399764+gym-cmd@users.noreply.github.com>
This commit is contained in:
parent
96c698030b
commit
55fe4a7894
12 changed files with 948 additions and 35 deletions
|
|
@ -163,7 +163,10 @@ from litellm.proxy.common_utils.user_api_key_cache import get_management_object_
|
|||
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
|
||||
id_jag_assertion_capture_gap_at_startup,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, get_server_root_path
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import (
|
||||
get_request_root_path,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.table_repositories import MCPServerRepository
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -3858,7 +3861,7 @@ class MCPServerManager:
|
|||
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
|
||||
# authorization_code's missing per-user token -> the per-server browser-OAuth
|
||||
# challenge, built here where the full MCPServer is in hand.
|
||||
raise_user_oauth_challenge(server, root_path=get_server_root_path())
|
||||
raise_user_oauth_challenge(server, root_path=get_request_root_path())
|
||||
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
|
||||
# token_exchange (OBO): a missing/rejected subject token -> the RFC 9728 challenge
|
||||
# pointing at the IdP the client must SSO with to obtain one, rather than an opaque
|
||||
|
|
@ -3866,7 +3869,7 @@ class MCPServerManager:
|
|||
# Access) threads its claims blob into the challenge for the client to satisfy.
|
||||
raise_token_exchange_challenge(
|
||||
server,
|
||||
root_path=get_server_root_path(),
|
||||
root_path=get_request_root_path(),
|
||||
claims=err.unauthorized.claims,
|
||||
)
|
||||
raise_public(err)
|
||||
|
|
@ -3914,7 +3917,7 @@ class MCPServerManager:
|
|||
if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
|
||||
return
|
||||
if subject_token is None and isinstance(spec.config, TokenExchangeConfig):
|
||||
raise_token_exchange_challenge(resolved_server, root_path=get_server_root_path())
|
||||
raise_token_exchange_challenge(resolved_server, root_path=get_request_root_path())
|
||||
match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
|
||||
case Ok(_):
|
||||
return
|
||||
|
|
@ -3922,7 +3925,7 @@ class MCPServerManager:
|
|||
if err.tag == "unauthorized" and isinstance(spec.config, TokenExchangeConfig):
|
||||
raise_token_exchange_challenge(
|
||||
resolved_server,
|
||||
root_path=get_server_root_path(),
|
||||
root_path=get_request_root_path(),
|
||||
claims=err.unauthorized.claims,
|
||||
)
|
||||
raise_public(err)
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ every other mode so the caller defers to v1 (parity-safe); it grows one branch p
|
|||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Final, Literal, NoReturn
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -310,13 +311,30 @@ def raise_public(error: CredError) -> NoReturn:
|
|||
def oauth_protected_resource_path(root_path: str, server: MCPServer) -> str:
|
||||
"""The server's RFC 9728 Protected Resource Metadata path, the shared anchor of both challenges.
|
||||
|
||||
``root_path`` is the proxy's ``SERVER_ROOT_PATH``, resolved by the caller (the imperative shell)
|
||||
so this stays a pure function of its inputs; ``"/"`` and ``""`` both mean no prefix. The path is
|
||||
relative, so it resolves against the caller's own host (correct even behind a reverse proxy).
|
||||
``root_path`` is the prefix the request was routed under, resolved by the caller (the imperative
|
||||
shell); ``"/"`` and ``""`` both mean no prefix. The path is relative, so it resolves against the
|
||||
caller's own host (correct even behind a reverse proxy).
|
||||
|
||||
URL structure depends on how the prefix is served:
|
||||
|
||||
- The scalar ``SERVER_ROOT_PATH`` deployment registers the well-known routes with the prefix
|
||||
*inserted* into the path (via :func:`well_known_root_suffix` at import time), matching RFC 8414
|
||||
§3 well-known path insertion. When ``root_path`` equals ``SERVER_ROOT_PATH`` the URL must use
|
||||
the same insertion or a client fetching it 404s.
|
||||
- The per-request ``SERVER_ROOT_PATHS`` deployment can't register routes per prefix (the prefix
|
||||
set is dynamic and could contain many entries); the middleware strips the prefix from
|
||||
``scope["path"]`` and the router matches the un-inserted well-known route. The URL must place
|
||||
the prefix *before* ``.well-known`` so the strip leaves a matching path.
|
||||
|
||||
Picking the wrong form 404s the client's discovery fetch — the discovery document and the 401
|
||||
challenge would then disagree on where the resource metadata lives.
|
||||
"""
|
||||
prefix: Final = "" if root_path == "/" else root_path
|
||||
name: Final = server.alias or server.server_name or server.name or server.server_id
|
||||
return f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}"
|
||||
scalar_env: Final = os.getenv("SERVER_ROOT_PATH", "").rstrip("/")
|
||||
if not prefix or (scalar_env and prefix == scalar_env):
|
||||
return f"/.well-known/oauth-protected-resource{prefix}/mcp/{name}"
|
||||
return f"{prefix}/.well-known/oauth-protected-resource/mcp/{name}"
|
||||
|
||||
|
||||
def raise_user_oauth_challenge(server: MCPServer, *, root_path: str) -> NoReturn:
|
||||
|
|
|
|||
|
|
@ -3848,12 +3848,14 @@ if MCP_AVAILABLE:
|
|||
# then exchanges. A tool-call-time 401 would be wrapped into a JSON-RPC error and the
|
||||
# header lost, so the discovery flow needs this pre-emptive challenge.
|
||||
if server and server.auth_type == MCPAuth.oauth2_token_exchange and not oauth2_headers:
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
from litellm.proxy.utils import get_server_root_path # noqa: PLC0415
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports proxy utils
|
||||
get_request_root_path,
|
||||
)
|
||||
|
||||
raise_token_exchange_challenge(server, root_path=get_server_root_path())
|
||||
raise_token_exchange_challenge(server, root_path=get_request_root_path())
|
||||
|
||||
# Exchange-backed modes (token_exchange's OBO mint, id_jag's stored-assertion mint): run
|
||||
# the exchange here at the transport edge, so a rejected subject raises the RFC 9728
|
||||
|
|
|
|||
|
|
@ -295,15 +295,18 @@ class LazyFeatureMiddleware:
|
|||
# Short-circuit once every feature has loaded.
|
||||
if scope["type"] in ("http", "websocket") and len(self._loaded) < len(self._features):
|
||||
path = scope.get("path", "")
|
||||
# Strip SERVER_ROOT_PATH so prefix matching works under a server
|
||||
# root path. Without this, requests like /api/v1/policies/... never
|
||||
# match the registered prefixes (/policies/...) and lazy features
|
||||
# stay unloaded — every endpoint under them returns 404. The
|
||||
# Strip the request's root_path so prefix matching works under a
|
||||
# server root path. Without this, requests like /api/v1/policies/...
|
||||
# never match the registered prefixes (/policies/...) and lazy
|
||||
# features stay unloaded — every endpoint under them returns 404.
|
||||
# scope["root_path"] wins over the cached env scalar: FastAPI
|
||||
# stamps SERVER_ROOT_PATH there, and PerRequestRootPathMiddleware
|
||||
# resolves SERVER_ROOT_PATHS prefixes there per request. The
|
||||
# `+ "/"` boundary prevents false-positive matches (e.g. /apiv2
|
||||
# against root /api). If the path doesn't start with the prefix
|
||||
# (e.g. a reverse proxy already stripped it), we leave it alone.
|
||||
if self._root_path and path.startswith(self._root_path + "/"):
|
||||
path = path[len(self._root_path) :]
|
||||
# against root /api); a pre-stripped path is left alone.
|
||||
root_path: Final = str(scope.get("root_path", "")).rstrip("/") or self._root_path
|
||||
if root_path and path.startswith(root_path + "/"):
|
||||
path = path[len(root_path) :] # rebind-ok: local strip after the boundary check above
|
||||
for feat in self._features:
|
||||
if feat.module_path in self._loaded:
|
||||
continue
|
||||
|
|
|
|||
122
litellm/proxy/middleware/per_request_root_path_middleware.py
Normal file
122
litellm/proxy/middleware/per_request_root_path_middleware.py
Normal file
|
|
@ -0,0 +1,122 @@
|
|||
"""Per-request ``root_path`` resolution from ``SERVER_ROOT_PATHS``.
|
||||
|
||||
``SERVER_ROOT_PATH`` is a startup scalar, so one deployment serves exactly one
|
||||
client-visible URL path prefix; a request under any other prefix 404s before a
|
||||
handler runs. When the ingress preserves several prefixes into one pod (e.g.
|
||||
``/tenant-a/*`` and ``/tenant-b/*``), the matched prefix becomes that
|
||||
request's ``scope["root_path"]`` instead: Starlette strips it during route
|
||||
matching and rebuilds it into ``request.base_url``, so every emitted URL —
|
||||
the MCP OAuth discovery ``resource`` (RFC 9728 §3) and the 401 challenges'
|
||||
``resource_metadata`` among them — lands under the prefix the client called.
|
||||
Opt-in: with ``SERVER_ROOT_PATHS`` unset the middleware is not added at all.
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Sequence
|
||||
from contextvars import ContextVar
|
||||
from typing import Final
|
||||
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
SERVER_ROOT_PATHS_ENV: Final = "SERVER_ROOT_PATHS"
|
||||
|
||||
# The effective ``root_path`` for the currently-handled request. Populated by
|
||||
# ``PerRequestRootPathMiddleware`` from the (possibly-mutated) scope so code
|
||||
# that emits URLs off the request path — the 401 challenges' resource_metadata
|
||||
# and ``get_custom_url``'s SSO callbacks among them — can pick up the prefix
|
||||
# the client actually called without threading scope through every call site.
|
||||
# ``None`` means "middleware did not run" (the ``SERVER_ROOT_PATHS`` env is
|
||||
# unset, so no per-request prefix exists); readers fall back to the scalar
|
||||
# ``SERVER_ROOT_PATH`` in that case, which matches the pre-middleware behavior.
|
||||
_request_root_path_var: Final[ContextVar[str | None]] = ContextVar("_request_root_path_var", default=None)
|
||||
|
||||
|
||||
def get_request_root_path() -> str:
|
||||
"""Return the effective ``root_path`` for the current request.
|
||||
|
||||
Reads the value ``PerRequestRootPathMiddleware`` stashed for this request;
|
||||
falls back through :func:`~litellm.proxy.utils.get_server_root_path` (i.e.
|
||||
the ``SERVER_ROOT_PATH`` env) when the middleware did not run — the
|
||||
scalar-only deployment. Delegating to the existing helper keeps every
|
||||
existing ``monkeypatch.setattr("litellm.proxy.utils.get_server_root_path"``
|
||||
test override working, and keeps a single source of truth for the scalar.
|
||||
"""
|
||||
value: Final = _request_root_path_var.get()
|
||||
if value is not None:
|
||||
return value
|
||||
# Lazy import: utils.py imports this module (via the lazy import inside
|
||||
# get_custom_url), so a top-level import would build a cycle at load time.
|
||||
from litellm.proxy.utils import get_server_root_path # noqa: PLC0415 # lazy import breaks a two-way dep
|
||||
|
||||
return get_server_root_path()
|
||||
|
||||
|
||||
def normalize_root_paths(raw_paths: Sequence[str]) -> tuple[str, ...]:
|
||||
"""Strip whitespace and trailing slashes, dedupe, order longest-first;
|
||||
warn and drop entries missing a leading ``/`` and the bare root."""
|
||||
kept: Final[list[str]] = [] # mutable-ok: local accumulator; escapes only as a tuple
|
||||
for entry in raw_paths:
|
||||
candidate = entry.strip()
|
||||
if not candidate:
|
||||
continue
|
||||
if not candidate.startswith("/"):
|
||||
verbose_proxy_logger.warning(
|
||||
"%s entry %r does not start with '/' and will be ignored.",
|
||||
SERVER_ROOT_PATHS_ENV,
|
||||
entry,
|
||||
)
|
||||
continue
|
||||
candidate = candidate.rstrip("/")
|
||||
if not candidate:
|
||||
verbose_proxy_logger.warning(
|
||||
"%s entry %r is the bare root and will be ignored; a root-mounted deployment needs no entry.",
|
||||
SERVER_ROOT_PATHS_ENV,
|
||||
entry,
|
||||
)
|
||||
continue
|
||||
if candidate not in kept:
|
||||
kept.append(candidate)
|
||||
return tuple(sorted(kept, key=len, reverse=True))
|
||||
|
||||
|
||||
def get_server_root_paths() -> tuple[str, ...]:
|
||||
"""The normalized ``SERVER_ROOT_PATHS`` prefixes, empty when unset."""
|
||||
configured: Final = os.getenv(SERVER_ROOT_PATHS_ENV, "")
|
||||
if not configured.strip():
|
||||
return ()
|
||||
return normalize_root_paths(configured.split(","))
|
||||
|
||||
|
||||
class PerRequestRootPathMiddleware:
|
||||
"""Sets ``scope["root_path"]`` to the configured prefix matching the
|
||||
request path on a whole-segment boundary. ``scope["path"]`` is left
|
||||
untouched (Starlette strips ``root_path`` at route-match time). Must be
|
||||
the outermost middleware so inner middlewares and the router see the
|
||||
resolved value; a matched prefix overrides a scalar ``SERVER_ROOT_PATH``
|
||||
for that request.
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp, root_paths: Sequence[str]) -> None:
|
||||
self.app = app
|
||||
self.root_paths: Final = normalize_root_paths(root_paths)
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] in ("http", "websocket"):
|
||||
path: Final = scope.get("path", "")
|
||||
for prefix in self.root_paths:
|
||||
if path == prefix or path.startswith(prefix + "/"):
|
||||
scope["root_path"] = prefix # rebind-ok: ASGI middleware contract; Router and base_url read it
|
||||
break
|
||||
# Stash the effective root_path (matched prefix, or the scope's
|
||||
# existing value when nothing matched — i.e. FastAPI's scalar
|
||||
# SERVER_ROOT_PATH) so code that emits URLs off the request path
|
||||
# picks the same prefix the router will resolve the request under.
|
||||
token: Final = _request_root_path_var.set(str(scope.get("root_path", "")))
|
||||
try:
|
||||
await self.app(scope, receive, send)
|
||||
finally:
|
||||
_request_root_path_var.reset(token)
|
||||
return
|
||||
await self.app(scope, receive, send)
|
||||
|
|
@ -598,6 +598,10 @@ from litellm.proxy.middleware.admission_control_middleware import (
|
|||
from litellm.proxy.middleware.in_flight_requests_middleware import (
|
||||
InFlightRequestsMiddleware,
|
||||
)
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import (
|
||||
PerRequestRootPathMiddleware,
|
||||
get_server_root_paths,
|
||||
)
|
||||
from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware
|
||||
from litellm.proxy.middleware.request_size_limit_middleware import (
|
||||
RequestSizeLimitMiddleware,
|
||||
|
|
@ -18384,6 +18388,22 @@ app.add_middleware(
|
|||
get_settings=lambda: get_admission_control_settings(general_settings),
|
||||
state=admission_control_state,
|
||||
)
|
||||
# Added last on purpose - last-added is outermost, and the client-visible URL
|
||||
# prefix must be resolved into scope["root_path"] before any inner middleware
|
||||
# or the router inspects the path. Only added when SERVER_ROOT_PATHS is
|
||||
# configured, so the default deployment's middleware stack is unchanged.
|
||||
_server_root_paths: Final = get_server_root_paths()
|
||||
if _server_root_paths:
|
||||
if server_root_path and server_root_path != "/":
|
||||
verbose_proxy_logger.warning(
|
||||
"Both SERVER_ROOT_PATH=%r and SERVER_ROOT_PATHS=%r are set. A request "
|
||||
"matching a SERVER_ROOT_PATHS prefix overrides the scalar root_path for "
|
||||
"that request; unmatched requests keep SERVER_ROOT_PATH. Configure one "
|
||||
"mechanism or the other.",
|
||||
server_root_path,
|
||||
_server_root_paths,
|
||||
)
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=_server_root_paths)
|
||||
|
||||
|
||||
async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "StreamingResponse":
|
||||
|
|
|
|||
|
|
@ -7212,7 +7212,19 @@ def get_custom_url(request_base_url: str, route: str | None = None) -> str:
|
|||
else:
|
||||
base_url = request_base_url
|
||||
|
||||
server_root_path: Final = get_server_root_path()
|
||||
# get_request_root_path() returns the prefix the router is actually
|
||||
# resolving this request under: the matched SERVER_ROOT_PATHS entry when
|
||||
# PerRequestRootPathMiddleware ran, otherwise the SERVER_ROOT_PATH scalar.
|
||||
# This keeps the emitted URL under one prefix — the one the client called —
|
||||
# instead of stacking the scalar onto a request already living under a
|
||||
# dynamic prefix (which would produce /tenant-a/legacy/... — a path that
|
||||
# doesn't exist). join_paths()'s tail-dedup then collapses the append when
|
||||
# base_url (i.e. request.base_url) already ends in the same prefix.
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import ( # noqa: PLC0415 # lazy: middleware imports utils
|
||||
get_request_root_path,
|
||||
)
|
||||
|
||||
server_root_path: Final = get_request_root_path()
|
||||
if route is not None:
|
||||
if server_root_path != "":
|
||||
# First join base_url with server_root_path, then with route
|
||||
|
|
|
|||
|
|
@ -476,16 +476,47 @@ def test_raise_public_plain_unauthorized_has_no_challenge():
|
|||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"root_path, expected_prefix",
|
||||
"root_path",
|
||||
[
|
||||
("/", ""), # "/" means no prefix
|
||||
("", ""), # empty means no prefix
|
||||
("/api/v1", "/api/v1"), # a real root path is prepended verbatim
|
||||
"/", # "/" means no prefix
|
||||
"", # empty means no prefix
|
||||
],
|
||||
)
|
||||
def test_oauth_protected_resource_path_honors_root_path(root_path, expected_prefix):
|
||||
def test_oauth_protected_resource_path_no_prefix(root_path, monkeypatch):
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
path = oauth_protected_resource_path(root_path, _server(alias="my-srv"))
|
||||
assert path == f"/.well-known/oauth-protected-resource{expected_prefix}/mcp/my-srv"
|
||||
assert path == "/.well-known/oauth-protected-resource/mcp/my-srv"
|
||||
|
||||
|
||||
def test_oauth_protected_resource_path_scalar_prefix_uses_rfc8414_insertion(monkeypatch):
|
||||
# A scalar SERVER_ROOT_PATH deployment registers the well-known routes with
|
||||
# the prefix inserted (via well_known_root_suffix at import time). The URL
|
||||
# must match that insertion or a client fetching it 404s.
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/api/v1")
|
||||
path = oauth_protected_resource_path("/api/v1", _server(alias="my-srv"))
|
||||
assert path == "/.well-known/oauth-protected-resource/api/v1/mcp/my-srv"
|
||||
|
||||
|
||||
def test_oauth_protected_resource_path_per_request_prefix_goes_before_wellknown(monkeypatch):
|
||||
# Per-request deployment: SERVER_ROOT_PATHS matched /tenant-a for this
|
||||
# request but the scalar SERVER_ROOT_PATH is unset. Routes were registered
|
||||
# without the well-known insertion, so the URL must place the prefix
|
||||
# *before* .well-known — PerRequestRootPathMiddleware strips it and the
|
||||
# router matches the un-inserted route.
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
path = oauth_protected_resource_path("/tenant-a", _server(alias="my-srv"))
|
||||
assert path == "/tenant-a/.well-known/oauth-protected-resource/mcp/my-srv"
|
||||
|
||||
|
||||
def test_oauth_protected_resource_path_dynamic_prefix_wins_over_scalar(monkeypatch):
|
||||
# Both env vars configured: the middleware matched a SERVER_ROOT_PATHS
|
||||
# prefix (/tenant-a) that differs from the scalar (/legacy). The URL must
|
||||
# advertise /tenant-a — the prefix the client called — with no /legacy
|
||||
# segment stacked onto it. Same review-fix invariant get_custom_url pins.
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
path = oauth_protected_resource_path("/tenant-a", _server(alias="my-srv"))
|
||||
assert path == "/tenant-a/.well-known/oauth-protected-resource/mcp/my-srv"
|
||||
assert "/legacy" not in path
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -510,7 +541,11 @@ def test_raise_user_oauth_challenge_points_at_per_server_prm():
|
|||
)
|
||||
|
||||
|
||||
def test_raise_user_oauth_challenge_includes_server_root_path():
|
||||
def test_raise_user_oauth_challenge_includes_server_root_path(monkeypatch):
|
||||
# The scalar deployment: routes are registered with the prefix inserted
|
||||
# (via well_known_root_suffix at import time), so the challenge URL uses
|
||||
# the RFC 8414 §3 insertion form.
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/api/v1")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_user_oauth_challenge(_server(alias="my-srv"), root_path="/api/v1")
|
||||
assert (
|
||||
|
|
@ -519,6 +554,20 @@ def test_raise_user_oauth_challenge_includes_server_root_path():
|
|||
)
|
||||
|
||||
|
||||
def test_raise_user_oauth_challenge_per_request_prefix_is_routable(monkeypatch):
|
||||
# Per-request deployment (SERVER_ROOT_PATHS matched /tenant-a): the
|
||||
# challenge URL must place /tenant-a before .well-known so the client's
|
||||
# discovery fetch routes through the same middleware strip the original
|
||||
# request went through. The scalar-inserted form would 404 here.
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_user_oauth_challenge(_server(alias="my-srv"), root_path="/tenant-a")
|
||||
assert (
|
||||
exc_info.value.headers["WWW-Authenticate"]
|
||||
== 'Bearer resource_metadata="/tenant-a/.well-known/oauth-protected-resource/mcp/my-srv"'
|
||||
)
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_is_rfc9728_invalid_token():
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
|
|
@ -535,17 +584,30 @@ def test_raise_token_exchange_challenge_is_rfc9728_invalid_token():
|
|||
assert "error_description=" in www
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_includes_server_root_path():
|
||||
def test_raise_token_exchange_challenge_includes_server_root_path(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/api/v1")
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_token_exchange_challenge(_server(alias="obo-srv"), root_path="/api/v1")
|
||||
www = exc_info.value.headers["WWW-Authenticate"]
|
||||
assert 'resource_metadata="/.well-known/oauth-protected-resource/api/v1/mcp/obo-srv"' in www
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_per_request_prefix_is_routable(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
)
|
||||
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
raise_token_exchange_challenge(_server(alias="obo-srv"), root_path="/tenant-a")
|
||||
www = exc_info.value.headers["WWW-Authenticate"]
|
||||
assert 'resource_metadata="/tenant-a/.well-known/oauth-protected-resource/mcp/obo-srv"' in www
|
||||
|
||||
|
||||
def test_raise_token_exchange_challenge_static_form_is_unchanged_without_step_up():
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_token_exchange_challenge,
|
||||
|
|
@ -589,9 +651,7 @@ def test_id_jag_client_secret_maps_to_config():
|
|||
# ID-JAG asserts the user's id_token; the access_token default maps to id_token.
|
||||
assert spec.config.subject_token_type == "urn:ietf:params:oauth:token-type:id_token"
|
||||
assert isinstance(spec.config.client_auth, ClientSecretAuth)
|
||||
assert spec.config.client_auth.client_secret.get_secret_value() == (
|
||||
"litellm-client-secret"
|
||||
)
|
||||
assert spec.config.client_auth.client_secret.get_secret_value() == ("litellm-client-secret")
|
||||
|
||||
|
||||
def test_id_jag_private_key_maps_to_private_key_jwt_auth():
|
||||
|
|
@ -617,9 +677,7 @@ def test_id_jag_private_key_wins_over_client_secret():
|
|||
|
||||
|
||||
def test_id_jag_honors_explicit_subject_token_type():
|
||||
spec = to_server_spec(
|
||||
_id_jag_server(subject_token_type="urn:ietf:params:oauth:token-type:saml2")
|
||||
)
|
||||
spec = to_server_spec(_id_jag_server(subject_token_type="urn:ietf:params:oauth:token-type:saml2"))
|
||||
assert spec is not None and isinstance(spec.config, IdJagConfig)
|
||||
assert spec.config.subject_token_type == "urn:ietf:params:oauth:token-type:saml2"
|
||||
|
||||
|
|
|
|||
|
|
@ -10342,6 +10342,230 @@ async def test_upstream_resource_sent_on_dcr_bridge_relay_authorize():
|
|||
assert query["client_id"] == ["caller-client"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Per-request root_path (SERVER_ROOT_PATHS / PerRequestRootPathMiddleware):
|
||||
# one app fronting several client-visible URL path prefixes, each prefix's
|
||||
# discovery documents emitting URLs under the prefix the client called
|
||||
# (RFC 9728 §3 exact-match). A scalar PROXY_BASE_URL / SERVER_ROOT_PATH can
|
||||
# encode at most one prefix per pod; these tests pin the N-prefix case.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _no_proxy_base_url(monkeypatch):
|
||||
"""Discovery must derive URLs from the request in these tests, so the
|
||||
scalar env overrides are cleared."""
|
||||
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _isolated_mcp_registry():
|
||||
"""Fixture-owned registry state: snapshot the shared registry, hand the
|
||||
test an empty one, restore afterwards so nothing leaks between cases."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
saved = dict(global_mcp_server_manager.registry)
|
||||
global_mcp_server_manager.registry.clear()
|
||||
try:
|
||||
yield global_mcp_server_manager.registry
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
global_mcp_server_manager.registry.update(saved)
|
||||
|
||||
|
||||
def _prefixed_discovery_client(prefixes):
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import (
|
||||
PerRequestRootPathMiddleware,
|
||||
)
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=prefixes)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
class TestPerRequestRootPathDiscovery:
|
||||
def test_prefixed_wellknown_not_routable_without_middleware(self, _isolated_mcp_registry):
|
||||
"""Control: on a plain app (the only shape a scalar root_path can
|
||||
express), a prefixed well-known request 404s before any discovery
|
||||
builder runs — the routing gap this feature exists to close."""
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
|
||||
|
||||
server = _create_oauth2_server(server_id="srv_a", name="server_a", server_name="server_a", alias="server_a")
|
||||
_isolated_mcp_registry[server.server_id] = server
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
client = TestClient(app)
|
||||
resp = client.get("/tenant-a/.well-known/oauth-protected-resource/mcp/server_a")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_two_prefixes_one_app_each_resource_matches_the_called_url(
|
||||
self, _no_proxy_base_url, _isolated_mcp_registry
|
||||
):
|
||||
"""The multi-origin case itself: two prefixes served by the same app,
|
||||
each per-server document's ``resource`` equal to the URL its client
|
||||
called — including the prefix."""
|
||||
for sid, name in (("srv_a", "server_a"), ("srv_b", "server_b")):
|
||||
_isolated_mcp_registry[sid] = _create_oauth2_server(server_id=sid, name=name, server_name=name, alias=name)
|
||||
client = _prefixed_discovery_client(["/tenant-a", "/tenant-b"])
|
||||
|
||||
resp_a = client.get("/tenant-a/.well-known/oauth-protected-resource/mcp/server_a")
|
||||
resp_b = client.get("/tenant-b/.well-known/oauth-protected-resource/mcp/server_b")
|
||||
|
||||
assert resp_a.status_code == 200
|
||||
assert resp_a.json()["resource"] == "http://testserver/tenant-a/mcp/server_a"
|
||||
assert resp_b.status_code == 200
|
||||
assert resp_b.json()["resource"] == "http://testserver/tenant-b/mcp/server_b"
|
||||
|
||||
# Every URL the document advertises stays under the request's
|
||||
# prefix, so it resolves on this same app.
|
||||
for auth_server in resp_a.json()["authorization_servers"]:
|
||||
assert auth_server.startswith("http://testserver/tenant-a/")
|
||||
|
||||
def test_unprefixed_requests_unchanged_on_the_same_app(self, _no_proxy_base_url, _isolated_mcp_registry):
|
||||
"""Backward compat on the very same app: a root request emits the
|
||||
document byte-identical to a deployment without the middleware."""
|
||||
server = _create_oauth2_server(server_id="srv_a", name="server_a", server_name="server_a", alias="server_a")
|
||||
_isolated_mcp_registry[server.server_id] = server
|
||||
client = _prefixed_discovery_client(["/tenant-a"])
|
||||
resp = client.get("/.well-known/oauth-protected-resource/mcp/server_a")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["resource"] == "http://testserver/mcp/server_a"
|
||||
|
||||
def test_unlisted_prefix_404s(self, _no_proxy_base_url):
|
||||
client = _prefixed_discovery_client(["/tenant-a"])
|
||||
assert client.get("/tenant-c/.well-known/oauth-protected-resource/mcp/server_a").status_code == 404
|
||||
|
||||
def test_aggregate_documents_and_as_endpoints_under_prefix(self, _no_proxy_base_url, _isolated_mcp_registry):
|
||||
"""Aggregate PRM/AS documents carry the prefix, and the advertised
|
||||
authorize endpoint actually resolves under it — the 404 trap that
|
||||
invalidated prefixing discovery URLs without per-request routing
|
||||
(#35226 review round 1)."""
|
||||
client = _prefixed_discovery_client(["/tenant-a"])
|
||||
|
||||
prm = client.get("/tenant-a/.well-known/oauth-protected-resource/mcp")
|
||||
asm = client.get("/tenant-a/.well-known/oauth-authorization-server/mcp")
|
||||
|
||||
assert prm.status_code == 200
|
||||
assert prm.json()["resource"] == "http://testserver/tenant-a/mcp"
|
||||
assert prm.json()["authorization_servers"] == ["http://testserver/tenant-a/mcp"]
|
||||
|
||||
assert asm.status_code == 200
|
||||
assert asm.json()["issuer"] == "http://testserver/tenant-a/mcp"
|
||||
assert asm.json()["authorization_endpoint"] == "http://testserver/tenant-a/authorize"
|
||||
|
||||
# The prefixed authorize URL routes to the real handler (not 404):
|
||||
# under per-request root_path the whole app is reachable per-prefix,
|
||||
# so discovery may advertise prefixed AS endpoints safely.
|
||||
assert client.get("/tenant-a/authorize").status_code != 404
|
||||
|
||||
def test_passthrough_challenge_metadata_url_carries_prefix(self, _no_proxy_base_url):
|
||||
"""The WWW-Authenticate resource_metadata URL a 401 advertises must
|
||||
land under the request's prefix, or the client is bounced to a
|
||||
document whose ``resource`` cannot match the URL it called."""
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
get_passthrough_resource_metadata_url,
|
||||
)
|
||||
|
||||
def _scope(path, root_path=None):
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": path,
|
||||
"headers": [],
|
||||
"scheme": "http",
|
||||
"server": ("testserver", 80),
|
||||
"client": ("1.2.3.4", 4444),
|
||||
}
|
||||
if root_path is not None:
|
||||
scope["root_path"] = root_path
|
||||
return scope
|
||||
|
||||
prefixed = get_passthrough_resource_metadata_url(
|
||||
scope=_scope("/tenant-a/mcp/github", root_path="/tenant-a"),
|
||||
server_name="github",
|
||||
)
|
||||
assert prefixed == "http://testserver/tenant-a/.well-known/oauth-protected-resource/mcp/github"
|
||||
|
||||
# Regression guard: no root_path → today's URL, unchanged.
|
||||
bare = get_passthrough_resource_metadata_url(
|
||||
scope=_scope("/mcp/github"),
|
||||
server_name="github",
|
||||
)
|
||||
assert bare == "http://testserver/.well-known/oauth-protected-resource/mcp/github"
|
||||
|
||||
def test_user_oauth_challenge_url_routes_and_resource_matches_client_url(
|
||||
self, _no_proxy_base_url, _isolated_mcp_registry
|
||||
):
|
||||
"""The reviewer's expected end-state, pinned end-to-end: an MCP endpoint
|
||||
raising ``raise_user_oauth_challenge`` under a per-request prefix must
|
||||
emit a resource_metadata URL the client can actually fetch, and the
|
||||
document it returns must carry the same prefix the client originally
|
||||
called. If either half breaks the client's discovery is dead."""
|
||||
import re
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
raise_user_oauth_challenge,
|
||||
)
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import (
|
||||
PerRequestRootPathMiddleware,
|
||||
get_request_root_path,
|
||||
)
|
||||
|
||||
server = _create_oauth2_server(
|
||||
server_id="srv_a", name="server_a", server_name="server_a", alias="server_a"
|
||||
)
|
||||
_isolated_mcp_registry[server.server_id] = server
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
|
||||
@app.post("/mcp/{name}")
|
||||
def _mcp(name: str, request: Request):
|
||||
try:
|
||||
raise_user_oauth_challenge(server, root_path=get_request_root_path())
|
||||
except HTTPException as exc:
|
||||
return {"www_authenticate": exc.headers["WWW-Authenticate"]}
|
||||
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=["/tenant-a", "/tenant-b"])
|
||||
client = TestClient(app)
|
||||
|
||||
for prefix, mcp_url in (
|
||||
("/tenant-a", "http://testserver/tenant-a/mcp/server_a"),
|
||||
("/tenant-b", "http://testserver/tenant-b/mcp/server_a"),
|
||||
("", "http://testserver/mcp/server_a"),
|
||||
):
|
||||
call = client.post(f"{prefix}/mcp/server_a")
|
||||
assert call.status_code == 200, call.text
|
||||
www = call.json()["www_authenticate"]
|
||||
match = re.search(r'resource_metadata="([^"]+)"', www)
|
||||
assert match, www
|
||||
discovery = client.get(match.group(1))
|
||||
# The challenge URL must route (a client that can't fetch it has
|
||||
# no way to reach the resource metadata).
|
||||
assert discovery.status_code == 200, (
|
||||
f"challenge URL {match.group(1)} for prefix {prefix!r} 404s; "
|
||||
"the client can't reach the resource metadata."
|
||||
)
|
||||
# And the doc's `resource` must equal the URL the client called
|
||||
# (RFC 9728 §3 exact match): a mismatch bounces a strict client.
|
||||
assert discovery.json()["resource"] == mcp_url, discovery.json()
|
||||
|
||||
|
||||
def _s256(verifier: str) -> str:
|
||||
return urlsafe_b64encode(hashlib.sha256(verifier.encode("ascii")).digest()).rstrip(b"=").decode("ascii")
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,302 @@
|
|||
"""Tests for PerRequestRootPathMiddleware (``SERVER_ROOT_PATHS``).
|
||||
|
||||
One deployment fronting several client-visible URL path prefixes: the matched
|
||||
prefix becomes that request's ``root_path``, so Starlette route matching and
|
||||
``request.base_url`` — and therefore every URL the proxy emits, the MCP OAuth
|
||||
discovery documents among them — resolve under the prefix the client called.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import (
|
||||
PerRequestRootPathMiddleware,
|
||||
get_request_root_path,
|
||||
get_server_root_paths,
|
||||
normalize_root_paths,
|
||||
)
|
||||
|
||||
|
||||
class TestNormalizeRootPaths:
|
||||
def test_strips_whitespace_and_trailing_slash(self):
|
||||
assert normalize_root_paths([" /tenant-a/ ", "/tenant-b"]) == (
|
||||
"/tenant-a",
|
||||
"/tenant-b",
|
||||
)
|
||||
|
||||
def test_drops_empty_entries(self):
|
||||
assert normalize_root_paths(["", " ", "/tenant-a"]) == ("/tenant-a",)
|
||||
|
||||
def test_drops_entries_without_leading_slash(self):
|
||||
# A typo'd entry must not silently match nothing at request time.
|
||||
assert normalize_root_paths(["tenant-a", "/tenant-b"]) == ("/tenant-b",)
|
||||
|
||||
def test_drops_bare_root(self):
|
||||
# "/" would turn every request into a root_path rewrite; a
|
||||
# root-mounted deployment needs no entry at all.
|
||||
assert normalize_root_paths(["/", "/tenant-a"]) == ("/tenant-a",)
|
||||
|
||||
def test_dedupes(self):
|
||||
assert normalize_root_paths(["/t", "/t/", " /t "]) == ("/t",)
|
||||
|
||||
def test_longest_first_for_nested_prefixes(self):
|
||||
# Longest-first ordering is what makes the most-specific nested
|
||||
# prefix win at match time.
|
||||
assert normalize_root_paths(["/t", "/t/deep"]) == ("/t/deep", "/t")
|
||||
|
||||
|
||||
class TestGetServerRootPaths:
|
||||
def test_unset_env_is_empty(self, monkeypatch):
|
||||
monkeypatch.delenv("SERVER_ROOT_PATHS", raising=False)
|
||||
assert get_server_root_paths() == ()
|
||||
|
||||
def test_empty_env_is_empty(self, monkeypatch):
|
||||
monkeypatch.setenv("SERVER_ROOT_PATHS", "")
|
||||
assert get_server_root_paths() == ()
|
||||
|
||||
def test_comma_separated_entries(self, monkeypatch):
|
||||
monkeypatch.setenv("SERVER_ROOT_PATHS", "/tenant-a, /tenant-b/")
|
||||
assert get_server_root_paths() == ("/tenant-a", "/tenant-b")
|
||||
|
||||
|
||||
def _capture_scope_middleware(root_paths):
|
||||
"""Middleware wired to a downstream that records the scope it received."""
|
||||
captured = {}
|
||||
|
||||
async def downstream(scope, receive, send):
|
||||
captured.update(scope)
|
||||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||||
await send({"type": "http.response.body", "body": b""})
|
||||
|
||||
return PerRequestRootPathMiddleware(downstream, root_paths=root_paths), captured
|
||||
|
||||
|
||||
async def _run(mw, scope):
|
||||
async def receive():
|
||||
return {"type": "http.request", "body": b"", "more_body": False}
|
||||
|
||||
async def send(message):
|
||||
pass
|
||||
|
||||
await mw(scope, receive, send)
|
||||
|
||||
|
||||
class TestPerRequestRootPathMiddleware:
|
||||
@pytest.mark.asyncio
|
||||
async def test_matched_prefix_becomes_root_path_path_untouched(self):
|
||||
# Starlette's router strips root_path from the (unmodified) path at
|
||||
# match time, so the middleware must NOT rewrite scope["path"].
|
||||
mw, captured = _capture_scope_middleware(["/tenant-a"])
|
||||
await _run(mw, {"type": "http", "path": "/tenant-a/mcp/x", "method": "GET", "headers": []})
|
||||
assert captured["root_path"] == "/tenant-a"
|
||||
assert captured["path"] == "/tenant-a/mcp/x"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exact_prefix_matches(self):
|
||||
mw, captured = _capture_scope_middleware(["/tenant-a"])
|
||||
await _run(mw, {"type": "http", "path": "/tenant-a", "method": "GET", "headers": []})
|
||||
assert captured["root_path"] == "/tenant-a"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_segment_boundary_prevents_sibling_match(self):
|
||||
# /tenant-ab must not match the /tenant-a prefix.
|
||||
mw, captured = _capture_scope_middleware(["/tenant-a"])
|
||||
await _run(mw, {"type": "http", "path": "/tenant-ab/mcp", "method": "GET", "headers": []})
|
||||
assert "root_path" not in captured
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unmatched_path_untouched(self):
|
||||
mw, captured = _capture_scope_middleware(["/tenant-a"])
|
||||
await _run(mw, {"type": "http", "path": "/chat/completions", "method": "GET", "headers": []})
|
||||
assert "root_path" not in captured
|
||||
assert captured["path"] == "/chat/completions"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_longest_nested_prefix_wins(self):
|
||||
mw, captured = _capture_scope_middleware(["/t", "/t/deep"])
|
||||
await _run(mw, {"type": "http", "path": "/t/deep/mcp", "method": "GET", "headers": []})
|
||||
assert captured["root_path"] == "/t/deep"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_matched_prefix_overrides_scalar_root_path(self):
|
||||
# FastAPI(root_path=SERVER_ROOT_PATH) stamps the scalar before the
|
||||
# middleware stack runs; a matched dynamic prefix wins for that
|
||||
# request (combining both mechanisms is warned about at startup).
|
||||
mw, captured = _capture_scope_middleware(["/tenant-a"])
|
||||
await _run(
|
||||
mw,
|
||||
{"type": "http", "path": "/tenant-a/mcp", "root_path": "/legacy", "method": "GET", "headers": []},
|
||||
)
|
||||
assert captured["root_path"] == "/tenant-a"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unmatched_request_keeps_scalar_root_path(self):
|
||||
mw, captured = _capture_scope_middleware(["/tenant-a"])
|
||||
await _run(
|
||||
mw,
|
||||
{"type": "http", "path": "/legacy/mcp", "root_path": "/legacy", "method": "GET", "headers": []},
|
||||
)
|
||||
assert captured["root_path"] == "/legacy"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_scope_matched(self):
|
||||
mw, captured = _capture_scope_middleware(["/tenant-a"])
|
||||
await _run(mw, {"type": "websocket", "path": "/tenant-a/ws", "headers": []})
|
||||
assert captured["root_path"] == "/tenant-a"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lifespan_scope_passes_through(self):
|
||||
called = {}
|
||||
|
||||
async def downstream(scope, receive, send):
|
||||
called["scope"] = scope
|
||||
|
||||
mw = PerRequestRootPathMiddleware(downstream, root_paths=["/tenant-a"])
|
||||
await _run(mw, {"type": "lifespan"})
|
||||
assert called["scope"] == {"type": "lifespan"}
|
||||
|
||||
|
||||
def _routed_client(prefixes):
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/where")
|
||||
def where(request: Request):
|
||||
return {
|
||||
"base_url": str(request.base_url),
|
||||
"root_path": request.scope.get("root_path", ""),
|
||||
}
|
||||
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=prefixes)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
class TestEndToEndRouting:
|
||||
def test_two_prefixes_route_on_one_app(self):
|
||||
# The property a scalar SERVER_ROOT_PATH cannot provide: two
|
||||
# client-visible prefixes served by the same app, each request
|
||||
# reconstructing its own base URL.
|
||||
client = _routed_client(["/tenant-a", "/tenant-b"])
|
||||
|
||||
resp_a = client.get("/tenant-a/where")
|
||||
resp_b = client.get("/tenant-b/where")
|
||||
|
||||
assert resp_a.status_code == 200
|
||||
assert resp_a.json() == {
|
||||
"base_url": "http://testserver/tenant-a/",
|
||||
"root_path": "/tenant-a",
|
||||
}
|
||||
assert resp_b.status_code == 200
|
||||
assert resp_b.json() == {
|
||||
"base_url": "http://testserver/tenant-b/",
|
||||
"root_path": "/tenant-b",
|
||||
}
|
||||
|
||||
def test_unprefixed_route_still_served(self):
|
||||
client = _routed_client(["/tenant-a"])
|
||||
resp = client.get("/where")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["base_url"] == "http://testserver/"
|
||||
|
||||
def test_unlisted_prefix_404s(self):
|
||||
client = _routed_client(["/tenant-a"])
|
||||
assert client.get("/tenant-c/where").status_code == 404
|
||||
|
||||
|
||||
class TestGetRequestRootPath:
|
||||
"""``get_request_root_path`` is the accessor that plumbs the middleware's
|
||||
resolved prefix to code that doesn't have scope in hand — the 401 challenge
|
||||
builders in ``mcp_server_manager`` / ``server`` and ``get_custom_url`` on the
|
||||
SSO callback path. Reading the SERVER_ROOT_PATH scalar there would emit URLs
|
||||
under a prefix the client didn't call, and stack a second prefix onto ones it
|
||||
did (the two review points this fixture pins)."""
|
||||
|
||||
def test_falls_back_to_server_root_path_env_outside_a_request(self, monkeypatch):
|
||||
# Outside a request the middleware's ContextVar is unset. The scalar
|
||||
# env still owns the answer, so pre-middleware call sites (module-load
|
||||
# UI URL builders, background tasks) behave exactly as they did.
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
assert get_request_root_path() == "/legacy"
|
||||
|
||||
def test_returns_empty_string_when_no_env_and_no_request(self, monkeypatch):
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
assert get_request_root_path() == ""
|
||||
|
||||
def test_returns_matched_prefix_inside_a_request(self, monkeypatch):
|
||||
# With both env vars set, a request matching a SERVER_ROOT_PATHS prefix
|
||||
# must see that prefix — not the SERVER_ROOT_PATH scalar — so the URL
|
||||
# it emits stays under the prefix the router will resolve it against.
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
|
||||
seen: list[str] = []
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/where")
|
||||
def where():
|
||||
seen.append(get_request_root_path())
|
||||
return {}
|
||||
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=["/tenant-a"])
|
||||
client = TestClient(app)
|
||||
|
||||
assert client.get("/tenant-a/where").status_code == 200
|
||||
assert seen == ["/tenant-a"]
|
||||
|
||||
def test_unmatched_request_falls_through_to_scope_scalar(self, monkeypatch):
|
||||
# A request the middleware saw but did not match keeps whatever
|
||||
# scope["root_path"] the app was mounted under (the scalar). The
|
||||
# ContextVar still reflects the effective per-request answer, so
|
||||
# emitted URLs and the router agree even on the fallback path.
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
seen: list[str] = []
|
||||
|
||||
app = FastAPI(root_path="/legacy")
|
||||
|
||||
@app.get("/where")
|
||||
def where():
|
||||
seen.append(get_request_root_path())
|
||||
return {}
|
||||
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=["/tenant-a"])
|
||||
client = TestClient(app)
|
||||
|
||||
assert client.get("/legacy/where").status_code == 200
|
||||
assert seen == ["/legacy"]
|
||||
|
||||
def test_each_request_sees_its_own_prefix(self, monkeypatch):
|
||||
# Two sequential requests through the same app must each see the
|
||||
# prefix they arrived under, so one tenant's client is never sent
|
||||
# the URL of another tenant's origin.
|
||||
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
|
||||
seen: list[tuple[str, str]] = []
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/where")
|
||||
def where(tag: str):
|
||||
seen.append((tag, get_request_root_path()))
|
||||
return {}
|
||||
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=["/tenant-a", "/tenant-b"])
|
||||
client = TestClient(app)
|
||||
assert client.get("/tenant-a/where?tag=a").status_code == 200
|
||||
assert client.get("/tenant-b/where?tag=b").status_code == 200
|
||||
assert seen == [("a", "/tenant-a"), ("b", "/tenant-b")]
|
||||
|
||||
def test_context_var_reset_after_request(self, monkeypatch):
|
||||
# A ContextVar left set after the request finishes would poison the
|
||||
# module-load-time callers that read it lazily (they'd think they were
|
||||
# inside a request under the last-seen prefix).
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/where")
|
||||
def where():
|
||||
return {"prefix": get_request_root_path()}
|
||||
|
||||
app.add_middleware(PerRequestRootPathMiddleware, root_paths=["/tenant-a"])
|
||||
client = TestClient(app)
|
||||
assert client.get("/tenant-a/where").json() == {"prefix": "/tenant-a"}
|
||||
# After the request finishes, the scalar-env fallback owns the answer
|
||||
# again — nothing was left stashed from the last request's scope.
|
||||
assert get_request_root_path() == "/legacy"
|
||||
|
|
@ -9224,6 +9224,82 @@ class TestLazyFeatureMiddleware:
|
|||
else:
|
||||
assert loads == [], f"{case}: feature must not load"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"env_root_path,scope_root_path,request_path,should_load,case",
|
||||
[
|
||||
# Per-request root_path (PerRequestRootPathMiddleware under
|
||||
# SERVER_ROOT_PATHS) with no scalar env: strip and match.
|
||||
("", "/tenant-a", "/tenant-a/dummy/x", True, "per-request root_path strip"),
|
||||
# scope root_path is authoritative over the cached env scalar.
|
||||
("/api/v1", "/tenant-a", "/tenant-a/dummy/x", True, "scope wins over env scalar"),
|
||||
# Boundary check still applies to the per-request value.
|
||||
("", "/tenant-a", "/tenant-ab/dummy/x", False, "boundary check on scope root_path"),
|
||||
# Empty scope root_path falls back to the env scalar.
|
||||
("/api/v1", "", "/api/v1/dummy/x", True, "empty scope falls back to env"),
|
||||
],
|
||||
)
|
||||
async def test_per_request_root_path_handling(
|
||||
self, monkeypatch, env_root_path, scope_root_path, request_path, should_load, case
|
||||
):
|
||||
"""
|
||||
``scope["root_path"]`` must be stripped before prefix matching when
|
||||
set — the scalar SERVER_ROOT_PATH lands there via
|
||||
``FastAPI(root_path=...)``, and PerRequestRootPathMiddleware
|
||||
(SERVER_ROOT_PATHS) resolves a per-request prefix there. Otherwise
|
||||
lazily-registered features — the MCP OAuth discovery router among
|
||||
them — stay unloaded under a client-visible prefix and 404.
|
||||
"""
|
||||
from fastapi import FastAPI
|
||||
|
||||
from litellm.proxy._lazy_features import (
|
||||
LazyFeature,
|
||||
LazyFeatureMiddleware,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", env_root_path)
|
||||
|
||||
loads = []
|
||||
|
||||
def fake_register(app, module):
|
||||
loads.append(getattr(module, "__name__", "?"))
|
||||
|
||||
feat = LazyFeature(
|
||||
name=f"dummy_prr_{case}",
|
||||
module_path="json",
|
||||
path_prefixes=("/dummy",),
|
||||
register_fn=fake_register,
|
||||
)
|
||||
|
||||
async def downstream(scope, receive, send):
|
||||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||||
await send({"type": "http.response.body", "body": b""})
|
||||
|
||||
target_app = FastAPI()
|
||||
mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,))
|
||||
|
||||
async def receive():
|
||||
return {"type": "http.request", "body": b"", "more_body": False}
|
||||
|
||||
async def send(message):
|
||||
pass
|
||||
|
||||
await mw(
|
||||
{
|
||||
"type": "http",
|
||||
"path": request_path,
|
||||
"root_path": scope_root_path,
|
||||
"method": "GET",
|
||||
"headers": [],
|
||||
},
|
||||
receive,
|
||||
send,
|
||||
)
|
||||
if should_load:
|
||||
assert loads == ["json"], f"{case}: expected feature to load"
|
||||
else:
|
||||
assert loads == [], f"{case}: feature must not load"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_first_requests_only_register_once(self):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -314,3 +314,76 @@ def test_normalize_route_for_root_path_error_path_when_route_not_under_root(
|
|||
_clear_url_env(monkeypatch)
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/proxy")
|
||||
assert normalize_route_for_root_path("/other/v1/chat") is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_custom_url under a per-request prefix (SERVER_ROOT_PATHS):
|
||||
# when a request lives under a dynamic prefix, request.base_url already
|
||||
# contains it. Appending the SERVER_ROOT_PATH scalar on top produced
|
||||
# double-prefixed SSO callback / login URLs (a path that doesn't exist on
|
||||
# the deployment). These pin that the emitted URL now stays under one
|
||||
# prefix — the one the request actually arrived on.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_get_custom_url_uses_per_request_prefix_when_middleware_ran(monkeypatch):
|
||||
"""The middleware stashes the effective per-request prefix in a ContextVar.
|
||||
``get_custom_url`` reads that in preference to the SERVER_ROOT_PATH scalar,
|
||||
and ``join_paths``'s tail-dedup collapses the append so a request whose
|
||||
``base_url`` already ends in ``/tenant-a`` does not become
|
||||
``/tenant-a/legacy/route``."""
|
||||
_clear_url_env(monkeypatch)
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import (
|
||||
_request_root_path_var,
|
||||
)
|
||||
|
||||
token = _request_root_path_var.set("/tenant-a")
|
||||
try:
|
||||
# request.base_url already carries the tenant prefix; the scalar
|
||||
# SERVER_ROOT_PATH must not be re-appended on top.
|
||||
result = get_custom_url(
|
||||
request_base_url="https://request.example.com/tenant-a/",
|
||||
route="/v1/chat",
|
||||
)
|
||||
finally:
|
||||
_request_root_path_var.reset(token)
|
||||
|
||||
assert result == "https://request.example.com/tenant-a/v1/chat"
|
||||
|
||||
|
||||
def test_get_custom_url_no_double_prefix_when_both_env_vars_configured(monkeypatch):
|
||||
"""Regression guard for the review point: with SERVER_ROOT_PATH also set
|
||||
(scalar-legacy) and the request matched by SERVER_ROOT_PATHS (per-request),
|
||||
the emitted URL is under one prefix — the per-request one — never both."""
|
||||
_clear_url_env(monkeypatch)
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
monkeypatch.setenv("SERVER_ROOT_PATHS", "/tenant-a")
|
||||
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import (
|
||||
_request_root_path_var,
|
||||
)
|
||||
|
||||
token = _request_root_path_var.set("/tenant-a")
|
||||
try:
|
||||
# No "/legacy" ever appears — the fix pins the reviewer's expected
|
||||
# behavior: only one prefix should apply per request.
|
||||
result = get_custom_url("https://api.example.com/tenant-a", "/sso/callback")
|
||||
finally:
|
||||
_request_root_path_var.reset(token)
|
||||
|
||||
assert result == "https://api.example.com/tenant-a/sso/callback"
|
||||
assert "/legacy" not in result
|
||||
|
||||
|
||||
def test_get_custom_url_scalar_only_still_stamps_root_path(monkeypatch):
|
||||
"""Pre-middleware deployments have not opted into SERVER_ROOT_PATHS at all;
|
||||
the ContextVar stays unset and the SERVER_ROOT_PATH scalar owns the answer
|
||||
— the behavior every existing scalar-only deployment relies on."""
|
||||
_clear_url_env(monkeypatch)
|
||||
monkeypatch.setenv("SERVER_ROOT_PATH", "/legacy")
|
||||
|
||||
result = get_custom_url("https://request.example.com", "/v1/chat")
|
||||
|
||||
assert result == "https://request.example.com/legacy/v1/chat"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue