feat(rust): add the MCP gateway (#43470)

Co-authored-by: Yujong Lee <yujong@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-27 19:37:41 -07:00 • committed by GitHub
parent 6e0926edde
commit f184ace25b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
45 changed files with 4754 additions and 27 deletions

159
litellm-rust/Cargo.lock generated
View file

@ -731,6 +731,7 @@ checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90"
dependencies = [
"axum-core",
"bytes",
"form_urlencoded",
"futures-util",
"http 1.4.2",
"http-body 1.1.0",
@ -747,6 +748,7 @@ dependencies = [
"serde_core",
"serde_json",
"serde_path_to_error",
"serde_urlencoded",
"sync_wrapper",
"tokio",
"tower",
@ -3755,6 +3757,7 @@ dependencies = [
"litellm-core",
"litellm-gateway-auth",
"litellm-gateway-inference",
"litellm-gateway-mcp",
"litellm-gateway-ui",
"litellm-http",
"litellm-llms",
@ -3764,7 +3767,9 @@ dependencies = [
"serde",
"serde_json",
"tempfile",
"thiserror 2.0.19",
"tokio",
"tokio-util",
"tower",
"tower-sessions-moka-store",
"tracing",
@ -3830,6 +3835,32 @@ dependencies = [
"tokio",
]
[[package]]
name = "litellm-gateway-mcp"
version = "0.1.0"
dependencies = [
"axum",
"base64 0.22.1",
"futures-util",
"http 1.4.2",
"litellm-auth-types",
"litellm-config",
"litellm-secrets",
"moka",
"rmcp",
"rstest",
"serde",
"serde_json",
"sha2 0.10.9",
"sse-stream",
"thiserror 2.0.19",
"tokio",
"tokio-util",
"tower",
"url",
"uuid",
]
[[package]]
name = "litellm-gateway-ui"
version = "0.1.0"
@ -3908,6 +3939,8 @@ dependencies = [
"litellm-core-utils",
"rcgen",
"reqwest 0.12.28",
"reqwest 0.13.5",
"rmcp",
"rstest",
"rustls 0.23.42",
"rustls-native-certs",
@ -4489,6 +4522,18 @@ dependencies = [
"version_check",
]
[[package]]
name = "nix"
version = "0.31.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cf20d2fde8ff38632c426f1165ed7436270b44f199fc55284c38276f9db47c3d"
dependencies = [
"bitflags 2.13.1",
"cfg-if",
"cfg_aliases",
"libc",
]
[[package]]
name = "nom"
version = "7.1.3"
@ -5037,6 +5082,20 @@ dependencies = [
"unicode-ident",
]
[[package]]
name = "process-wrap"
version = "10.0.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1f21b97672d2dc848e7b25701ab4535618b92f4861c13cc3f7f7bed52ad3c8da"
dependencies = [
"futures",
"indexmap 2.14.0",
"nix",
"tokio",
"tracing",
"windows",
]
[[package]]
name = "proptest"
version = "1.11.0"
@ -5621,6 +5680,7 @@ dependencies = [
"bytes",
"futures-core",
"futures-util",
"h2 0.4.15",
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
@ -5676,6 +5736,40 @@ dependencies = [
"windows-sys 0.52.0",
]
[[package]]
name = "rmcp"
version = "3.4.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5b6317cd8c13e3ec9033cf2aa5aa92cd743f0f4f8a93cddc42ea0ae6ce3b8898"
dependencies = [
"async-trait",
"base64 0.23.1",
"bytes",
"chrono",
"futures",
"http 1.4.2",
"http-body 1.1.0",
"http-body-util",
"indexmap 2.14.0",
"pastey",
"pin-project-lite",
"process-wrap",
"rand 0.10.2",
"reqwest 0.13.5",
"schemars 1.2.2",
"serde",
"serde_json",
"sse-stream",
"thiserror 2.0.19",
"tokio",
"tokio-stream",
"tokio-util",
"tower-service",
"tracing",
"url",
"uuid",
]
[[package]]
name = "rsa"
version = "0.9.10"
@ -6000,6 +6094,7 @@ version = "1.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "687274d293b6cdc6e73e0fee520bf2049650090d7164f87672d212a3c530cf4a"
dependencies = [
"chrono",
"dyn-clone",
"ref-cast",
"schemars_derive",
@ -6610,6 +6705,19 @@ dependencies = [
"url",
]
[[package]]
name = "sse-stream"
version = "0.2.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c25ac7aff0abd1dbc474536e40416e1102c7dd9bfba0b9861c6d357f835dcfb4"
dependencies = [
"bytes",
"futures-util",
"http-body 1.1.0",
"http-body-util",
"pin-project-lite",
]
[[package]]
name = "stable_deref_trait"
version = "1.2.1"
@ -7911,6 +8019,27 @@ version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f"
[[package]]
name = "windows"
version = "0.62.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "527fadee13e0c05939a6a05d5bd6eec6cd2e3dbd648b9f8e447c6518133d8580"
dependencies = [
"windows-collections",
"windows-core",
"windows-future",
"windows-numerics",
]
[[package]]
name = "windows-collections"
version = "0.3.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "23b2d95af1a8a14a3c7367e1ed4fc9c20e0a26e79551b1454d72583c97cc6610"
dependencies = [
"windows-core",
]
[[package]]
name = "windows-core"
version = "0.62.2"
@ -7924,6 +8053,17 @@ dependencies = [
"windows-strings",
]
[[package]]
name = "windows-future"
version = "0.3.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e1d6f90251fe18a279739e78025bd6ddc52a7e22f921070ccdc67dde84c605cb"
dependencies = [
"windows-core",
"windows-link",
"windows-threading",
]
[[package]]
name = "windows-implement"
version = "0.60.2"
@ -7952,6 +8092,16 @@ version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
[[package]]
name = "windows-numerics"
version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6e2e40844ac143cdb44aead537bbf727de9b044e107a0f1220392177d15b0f26"
dependencies = [
"windows-core",
"windows-link",
]
[[package]]
name = "windows-result"
version = "0.4.1"
@ -8004,6 +8154,15 @@ dependencies = [
"windows_x86_64_msvc",
]
[[package]]
name = "windows-threading"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3949bd5b99cafdf1c7ca86b43ca564028dfe27d66958f2470940f73d86d75b37"
dependencies = [
"windows-link",
]
[[package]]
name = "windows_aarch64_gnullvm"
version = "0.52.6"

View file

@ -13,6 +13,7 @@ litellm-config = { path = "crates/config" }
litellm-router = { path = "crates/router" }
litellm-tracing = { path = "crates/tracing" }
litellm-core = { path = "crates/core" }
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
litellm-gateway = { path = "crates/gateway" }
litellm-gateway-inference = { path = "crates/gateway-inference" }
litellm-gateway-auth = { path = "crates/gateway-auth" }

View file

@ -1,4 +1,4 @@
# The Tokio runtime is reached only through `host-python/src/execution.rs`, whose fork gate
# The Tokio runtime is reached only through `host-python/src/runtime.rs`, whose fork gate
# must see every entry. Going around it makes a fork-after-use hang instead of raising.
disallowed-methods = [
{ path = "pyo3_async_runtimes::tokio::get_runtime", reason = "use litellm_host_python::run_sync / run_sync_value" },

View file

@ -1,5 +1,6 @@
mod error;
mod includes;
mod mcp;
mod model;
mod settings;
mod value;
@ -9,6 +10,7 @@ use std::{fmt, path::Path};
use serde::Deserialize;
pub use error::Error;
pub use mcp::{McpAuth, McpServer, McpTransport};
pub use model::{LiteLlmParams, Model};
pub use settings::{GeneralSettings, LiteLlmSettings, RouterSettings};
pub use value::{AdditionalFields, Flag, NumberOrString, Object, OneOrMany, Value};
@ -24,6 +26,7 @@ pub struct Config {
pub callback_settings: Object,
pub assistant_settings: Object,
pub default_vertex_config: Object,
pub mcp_servers: std::collections::BTreeMap<String, McpServer>,
pub credential_list: Box<[Object]>,
pub guardrails: Box<[Object]>,
pub prompts: Box<[Object]>,
@ -55,6 +58,7 @@ impl fmt::Debug for Config {
.field("callback_settings", &self.callback_settings)
.field("assistant_settings", &self.assistant_settings)
.field("default_vertex_config", &self.default_vertex_config)
.field("mcp_servers", &self.mcp_servers)
.field("credential_list", &self.credential_list)
.field("guardrails", &self.guardrails)
.field("prompts", &self.prompts)

View file

@ -0,0 +1,97 @@
use std::{collections::BTreeMap, fmt};
use litellm_auth_types::SecretValue;
use serde::Deserialize;
use crate::Object;
#[derive(Clone, Default, Deserialize)]
#[serde(default)]
pub struct McpServer {
pub server_id: Option<String>,
pub alias: Option<String>,
pub description: Option<String>,
pub mcp_info: Object,
pub transport: McpTransport,
pub url: Option<SecretValue>,
pub command: Option<String>,
pub args: Box<[String]>,
pub env: BTreeMap<String, SecretValue>,
pub auth_type: Option<McpAuth>,
#[serde(alias = "auth_value")]
pub authentication_token: Option<SecretValue>,
pub static_headers: BTreeMap<String, SecretValue>,
pub upstream_token_header: Option<String>,
pub allowed_tools: Option<Box<[String]>>,
pub timeout: Option<f64>,
pub max_concurrent_requests: Option<usize>,
#[serde(flatten)]
pub unsupported: Object,
}
impl fmt::Debug for McpServer {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("McpServer")
.field("transport", &self.transport)
.field("auth_type", &self.auth_type)
.field("timeout", &self.timeout)
.field("max_concurrent_requests", &self.max_concurrent_requests)
.finish_non_exhaustive()
}
}
#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum McpTransport {
#[default]
Http,
Sse,
Stdio,
}
impl McpTransport {
pub fn as_str(self) -> &'static str {
match self {
Self::Http => "http",
Self::Sse => "sse",
Self::Stdio => "stdio",
}
}
}
#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum McpAuth {
None,
ApiKey,
BearerToken,
Basic,
Authorization,
Token,
Oauth2,
AwsSigv4,
Oauth2TokenExchange,
Oauth2IdJag,
TruePassthrough,
OauthDelegate,
}
impl McpAuth {
pub fn as_str(self) -> &'static str {
match self {
Self::None => "none",
Self::ApiKey => "api_key",
Self::BearerToken => "bearer_token",
Self::Basic => "basic",
Self::Authorization => "authorization",
Self::Token => "token",
Self::Oauth2 => "oauth2",
Self::AwsSigv4 => "aws_sigv4",
Self::Oauth2TokenExchange => "oauth2_token_exchange",
Self::Oauth2IdJag => "oauth2_id_jag",
Self::TruePassthrough => "true_passthrough",
Self::OauthDelegate => "oauth_delegate",
}
}
}

View file

@ -35,6 +35,8 @@ pub struct GeneralSettings {
pub dangerously_permit_weak_or_unset_master_key: Option<bool>,
pub plugins: Option<Box<[Object]>>,
pub coordination_redis: Option<Object>,
pub mcp_allowed_hosts: Option<Box<[String]>>,
pub mcp_allowed_origins: Box<[String]>,
#[serde(flatten)]
pub additional_fields: AdditionalFields,
}
@ -69,6 +71,8 @@ impl Default for GeneralSettings {
dangerously_permit_weak_or_unset_master_key: None,
plugins: None,
coordination_redis: None,
mcp_allowed_hosts: None,
mcp_allowed_origins: Box::default(),
additional_fields: AdditionalFields::new(),
}
}

View file

@ -233,6 +233,9 @@ finetune_settings:
- custom_llm_provider: openai
mcp_tools:
- name: lookup
mcp_servers:
docs:
url: https://example.test/mcp
vector_store_registry:
- vector_store_name: docs
worker_registry:
@ -262,6 +265,7 @@ include:
assert_eq!(config.files_settings.len(), 1);
assert_eq!(config.finetune_settings.len(), 1);
assert_eq!(config.mcp_tools.len(), 1);
assert_eq!(config.mcp_servers.len(), 1);
assert_eq!(config.vector_store_registry.len(), 1);
assert_eq!(config.worker_registry.len(), 1);
assert_eq!(config.agents.len(), 1);
@ -341,3 +345,31 @@ fn resolves_nested_includes_once_in_breadth_first_order() {
);
assert!(config.include.is_empty());
}
#[rstest]
fn mcp_config_redacts_nested_credentials_and_preserves_policy_for_validation() {
let config = Config::from_yaml("mcp_servers:\n docs:\n url: https://example.test/private-secret/mcp\n authentication_token: upstream-secret\n static_headers: {x-token: header-secret}\n env: {TOKEN: env-secret}\n args: [argument-secret]\n client_secret: oauth-secret\n allowed_tools: [search]\n").unwrap();
let server = &config.mcp_servers["docs"];
assert_eq!(server.allowed_tools.as_deref().unwrap(), ["search"]);
assert!(server.unsupported.contains_key("client_secret"));
let debug = format!("{config:?}");
for secret in [
"private-secret",
"upstream-secret",
"header-secret",
"env-secret",
"oauth-secret",
"argument-secret",
] {
assert!(!debug.contains(secret));
}
}
#[rstest]
#[case::transport("transport: invalid")]
#[case::auth("auth_type: invalid")]
#[case::concurrency("max_concurrent_requests: -1")]
#[case::headers("static_headers: {x-token: [not, a, string]}")]
fn rejects_invalid_typed_mcp_settings(#[case] setting: &str) {
assert!(Config::from_yaml(&format!("mcp_servers:\n docs:\n {setting}\n")).is_err());
}

View file

@ -14,11 +14,13 @@ Inbound authentication uses a verifier, identity resolver, and authorizer inject
The Axum `authenticate` middleware currently accepts one Authorization header using the existing Bearer format. Missing, duplicate, empty, or malformed credentials fail. Authentication replaces any preexisting caller extension and checks the method and matched route before dispatch. Handlers extract `AuthenticatedRequest` and authorize their parsed operation before calling a provider. Missing authenticated context fails closed
The gateway also inserts an MCP `Authorization` extension and a `SessionOwner` derived from the verified principal and credential reference. The MCP host checks the injected policy at initialization and before sending resolved upstream operations. Standalone MCP library hosts retain their existing explicitly trusted registry behavior when no policy is installed; production mounts must install authentication and the policy extension together
Local UI login continues using axum-login and tower-sessions. Only after session and CSRF validation does `UiSession` expose an authenticated caller. Its credentials are restricted to session-info and logout operations, so the UI CSRF bearer cannot authorize inference. Future UI management routes must extend that explicit scope
Authentication evidence and principals contain no raw token, password, request body, or mutable accounting state. Session ownership is scoped by principal authority, subject, verifier, and credential ID, so separate credentials do not silently share an MCP session. Credential rotation may retain ownership when the verifier preserves a stable credential ID. Scope and expiry checks still run for each operation
Failures distinguish invalid or expired credentials, forbidden operations, unavailable authentication services, and missing server configuration/context. HTTP adapters map those outcomes to status codes; inference keeps its API-specific error envelopes
Failures distinguish invalid or expired credentials, forbidden operations, unavailable authentication services, and missing server configuration/context. HTTP adapters map those outcomes to status codes; inference keeps its API-specific error envelopes and MCP translates policy failures to its own protocol
Virtual-key storage, JWT verification, OAuth2 introspection, trusted-proxy validation, SSO, and custom Python hook adapters are not implemented here yet. They should supply these contracts rather than bypassing the shared authorization boundary. Request-body-dependent custom hooks will need a bounded endpoint adapter after parsing

View file

@ -23,7 +23,7 @@ use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use litellm_secrets::source::SecretSource;
pub use error::Error;
pub use litellm_router::{Deployment, Router as ModelList};
pub use litellm_router::{Deployment, Router as ModelRouter};
pub use request::{JsonObject, RequestId};
pub struct Gateway {
@ -32,7 +32,7 @@ pub struct Gateway {
pub messages: MessagesRoute,
pub ocr: OcrRoute,
pub responses: ResponsesRoute,
pub models: ModelList,
pub models: ModelRouter,
pub secrets: Arc<dyn SecretSource>,
pub resources: CoreResources,
pub http: HttpClientConfig,
@ -43,7 +43,7 @@ impl Gateway {
resources: CoreResources,
http: HttpClientConfig,
secrets: Arc<dyn SecretSource>,
models: ModelList,
models: ModelRouter,
) -> Result<Self, litellm_http::Error> {
let provider = resources.pool.client(&http, ClientVariant::Provider)?;
let auth = resources.auth.clone();

View file

@ -0,0 +1,50 @@
# MCP gateway
Native MCP routing using the official [Rust SDK](https://github.com/modelcontextprotocol/rust-sdk), pinned to `rmcp 3.4.1`. The crate exposes a mountable Axum router and an SDK `ServerHandler`
The HTTP router serves `/mcp`, `/{server}/mcp`, `/mcp/{server}`, `/mcp/{server}/mcp`, legacy `/mcp/sse` and `/mcp/sse/messages`, `/mcp/enabled`, and `/mcp-rest/tools/list` and `/mcp-rest/tools/call`. Streamable HTTP supports stateless requests and legacy initialize/session flows. The caller supplies the shutdown token, allowed hosts, allowed origins, and server identity through `HttpConfig`
`NativeGateway` aggregates upstream tools, prompts, resources, and resource templates. Tool and prompt calls resolve registered server prefixes, including prefixes containing hyphens. Resource names receive prefixes while their URIs remain unchanged, and reads require exactly one selected server. Tool allowlists apply to discovery and execution. Upstream pagination is drained with cycle detection. Aggregate tool listings retain healthy servers and expose sanitized outcomes under `_meta["litellm.ai/server_outcomes"]`, while scoped REST listings preserve upstream failures
## Integration
Construct `Server` entries from connected SDK peers and pass a `Registry` to `NativeGateway::new`. Keep their `RunningService` handles alive in the host and cancel them during shutdown. SDK peers can use stdio, Streamable HTTP, or another SDK transport. Connection establishment, outbound credentials, HTTP client pools, and reconnection remain owned by the host
Mount `router(operations, config)` behind the gateway's authentication middleware before merging it with other endpoint routers. `Registry` is an authorized static catalog, not an authenticator. Production hosts must authenticate all requests, including initialize, SSE GET, session POST, and DELETE. Authentication middleware may insert `SessionOwner` into request extensions to bind sessions to a stable principal across credential rotation. Otherwise sessions bind to a digest of the presented credentials and requested server scope
Implement `ServerResolver` to select authorized peers and tool permissions for each request. It receives HTTP headers and trusted extensions through `Context.parts`, plus the downstream MCP request context when available. Path and `x-mcp-servers` scopes can narrow that selection, never broaden it. Outbound headers are not copied automatically
Wrap `Operations` to integrate guardrails, spend logging, virtual tools, and other policy. Both MCP and REST execution use this interface. Database management, OAuth/BYOK services, toolsets, OpenAPI conversion, and semantic search are integration responsibilities in this stage, rather than new implementations in this crate
`RelayClient::for_request` forwards sampling, elicitation, roots, and progress to the originating downstream client. Create it for that request's upstream connection and retain the connection until execution completes. Do not share this relay across callers. Sampling can instead be handled by an injected SDK `ClientHandler` backed by inference
## Local interoperability check
From the repository root, this loopback-only example runs the existing Python MCP fixture through the SDK's stdio transport. It has no admission authentication and is intended for local testing
```sh
cargo run --manifest-path litellm-rust/Cargo.toml -p litellm-gateway-mcp --example stdio_proxy -- \
.venv/bin/python -c 'import runpy; runpy.run_path("tests/mcp_tests/mcp_e2e_upstream_server.py")["mcp"].run(transport="stdio")'
```
```sh
curl -sS http://127.0.0.1:4000/mcp-rest/tools/call \
-H 'Content-Type: application/json' \
-d '{"server_id":"local","name":"add","arguments":{"a":2,"b":3}}'
```
```json
{"content":[{"type":"text","text":"5"}],"structuredContent":{"result":5},"isError":false}
```
```sh
curl -sS http://127.0.0.1:4000/mcp \
-H 'Content-Type: application/json' \
-H 'Accept: application/json, text/event-stream' \
-H 'MCP-Protocol-Version: 2025-11-25' \
-d '{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"local-multiply","arguments":{"a":6,"b":7}}}'
```
```text
data: {"jsonrpc":"2.0","id":1,"result":{"content":[{"type":"text","text":"42"}],"structuredContent":{"result":42},"isError":false}}
```

View file

@ -0,0 +1,33 @@
[package]
name = "litellm-gateway-mcp"
version = "0.1.0"
edition.workspace = true
rust-version.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
axum = { workspace = true, features = ["json", "query"] }
futures-util.workspace = true
litellm-config.workspace = true
litellm-secrets.workspace = true
litellm-auth-types.workspace = true
base64.workspace = true
http.workspace = true
uuid.workspace = true
url.workspace = true
moka.workspace = true
sha2.workspace = true
tower = { version = "0.5", features = ["util"] }
rmcp = { version = "=3.4.1", default-features = false, features = ["server", "client", "elicitation", "transport-streamable-http-server", "transport-streamable-http-client", "transport-child-process"] }
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tokio.workspace = true
tokio-util = "0.7"
[dev-dependencies]
rstest.workspace = true
tokio = { workspace = true, features = ["io-util", "signal"] }
rmcp = { version = "=3.4.1", default-features = false, features = ["transport-child-process"] }
sse-stream = "0.2.6"

View file

@ -0,0 +1,40 @@
use std::{error::Error, sync::Arc, time::Duration};
use litellm_gateway_mcp::{HttpConfig, NativeGateway, Registry, Server, ServerInfo, router};
use rmcp::{ServiceExt, transport::TokioChildProcess};
#[tokio::main]
async fn main() -> Result<(), Box<dyn Error>> {
let arguments: Vec<_> = std::env::args_os().skip(1).collect();
let (program, arguments) = arguments
.split_first()
.ok_or("usage: stdio_proxy <program> [arguments...]")?;
let mut command = tokio::process::Command::new(program);
command.args(arguments).kill_on_drop(true);
let upstream = ().serve(TokioChildProcess::new(command)?).await?;
let server = Server {
info: ServerInfo {
server_id: "local".into(),
server_name: "local".into(),
alias: None,
},
peer: upstream.peer().clone(),
allowed_tools: None,
};
let gateway = Arc::new(NativeGateway::new(
Arc::new(Registry::new(vec![server])?),
Duration::from_secs(30),
));
let config = HttpConfig::default();
let shutdown = config.cancellation_token.clone();
let listener = tokio::net::TcpListener::bind("127.0.0.1:4000").await?;
eprintln!("Local MCP example listening on http://127.0.0.1:4000/mcp");
axum::serve(listener, router(gateway, config))
.with_graceful_shutdown(async move {
let _ = tokio::signal::ctrl_c().await;
shutdown.cancel();
})
.await?;
upstream.cancel().await?;
Ok(())
}

View file

@ -0,0 +1,477 @@
use std::{
collections::{BTreeMap, HashMap},
num::NonZeroUsize,
sync::Arc,
time::Duration,
};
use base64::{Engine, engine::general_purpose::STANDARD};
use futures_util::future::try_join_all;
use http::{HeaderName, HeaderValue};
use litellm_auth_types::SecretValue;
use litellm_config::{McpAuth, McpServer, McpTransport};
use litellm_secrets::source::SecretSource;
use rmcp::{
RoleClient, ServiceExt,
service::RunningService,
transport::{
StreamableHttpClientTransport, TokioChildProcess,
streamable_http_client::{StreamableHttpClient, StreamableHttpClientTransportConfig},
},
};
use sha2::{Digest, Sha256};
use tokio_util::sync::CancellationToken;
use crate::{ConnectError, Limits, NativeGateway, Registry, Server, ServerInfo};
pub struct ConfiguredGateway {
pub operations: Arc<NativeGateway>,
services: Vec<RunningService<RoleClient, ()>>,
}
fn invalid(server: &str, message: impl Into<String>) -> ConnectError {
ConnectError::Configuration {
server: server.into(),
message: message.into(),
}
}
fn timeout(name: &str, config: &McpServer) -> Result<Duration, ConnectError> {
let duration = Duration::try_from_secs_f64(config.timeout.unwrap_or(30.0))
.map_err(|_| invalid(name, "timeout must be finite and positive"))?;
if duration.is_zero() {
return Err(invalid(name, "timeout must be positive"));
}
Ok(duration)
}
fn info(name: &str, config: &McpServer) -> ServerInfo {
let identity = format!(
"{name}|{}|{}|{}|{}",
config.url.as_ref().map_or("", SecretValue::expose),
config.transport.as_str(),
config.auth_type.map_or("", McpAuth::as_str),
config.alias.as_deref().unwrap_or("")
);
ServerInfo {
server_id: config
.server_id
.clone()
.unwrap_or_else(|| format!("{:x}", Sha256::digest(identity))[..32].into()),
server_name: name.into(),
alias: config.alias.clone(),
}
}
fn validate(name: &str, config: &McpServer) -> Result<(), ConnectError> {
timeout(name, config)?;
if name.trim().is_empty()
|| config
.server_id
.as_ref()
.is_some_and(|id| id.trim().is_empty())
|| config
.alias
.as_ref()
.is_some_and(|alias| alias.trim().is_empty())
{
return Err(invalid(name, "server identifiers must not be blank"));
}
if let Some(field) = config.unsupported.keys().next() {
return Err(invalid(
name,
format!("setting '{field}' is not supported by the Rust MCP host"),
));
}
if config
.max_concurrent_requests
.is_some_and(|limit| limit == 0 || limit > tokio::sync::Semaphore::MAX_PERMITS)
{
return Err(invalid(name, "max_concurrent_requests is out of range"));
}
if config.mcp_info.contains_key("protocol_version")
|| config.mcp_info.contains_key("mcp_server_cost_info")
{
return Err(invalid(
name,
"mcp_info protocol and cost settings are not supported",
));
}
if !matches!(
config.auth_type,
None | Some(
McpAuth::None
| McpAuth::ApiKey
| McpAuth::BearerToken
| McpAuth::Basic
| McpAuth::Authorization
| McpAuth::Token
)
) {
return Err(invalid(
name,
"configured upstream auth mode is not supported by the Rust MCP host",
));
}
match config.transport {
McpTransport::Sse => {
return Err(invalid(
name,
"legacy SSE upstreams are not supported by the pinned SDK; use transport: http",
));
}
McpTransport::Http => {
if config.url.is_none()
|| config.command.is_some()
|| !config.args.is_empty()
|| !config.env.is_empty()
{
return Err(invalid(
name,
"HTTP transport requires url and does not accept command, args or env",
));
}
}
McpTransport::Stdio => {
if config
.command
.as_ref()
.is_none_or(|command| command.trim().is_empty())
|| config.url.is_some()
{
return Err(invalid(
name,
"stdio transport requires command and does not accept url",
));
}
if config.authentication_token.is_some()
|| !config.static_headers.is_empty()
|| config.upstream_token_header.is_some()
|| !matches!(config.auth_type, None | Some(McpAuth::None))
{
return Err(invalid(
name,
"stdio credentials must be supplied through env",
));
}
}
}
if matches!(config.auth_type, None | Some(McpAuth::None)) {
if config.authentication_token.is_some() || config.upstream_token_header.is_some() {
return Err(invalid(
name,
"authentication_token and upstream_token_header require an auth_type",
));
}
} else if config.authentication_token.is_none() {
return Err(invalid(
name,
"configured auth_type requires authentication_token",
));
}
Ok(())
}
async fn resolve(value: &SecretValue, secrets: &dyn SecretSource) -> Result<String, ConnectError> {
match value.expose().strip_prefix("os.environ/") {
Some(key) => secrets
.get_secret_str(key)
.await?
.map(|secret| secret.expose().to_owned())
.ok_or_else(|| invalid("configuration", "referenced secret is unavailable")),
None => Ok(value.expose().to_owned()),
}
}
fn header(name: &str, value: &str) -> Result<(HeaderName, HeaderValue), ConnectError> {
let key = HeaderName::try_from(name)
.map_err(|_| invalid("configuration", "invalid upstream header name"))?;
if matches!(
key.as_str(),
"host"
| "content-length"
| "content-type"
| "transfer-encoding"
| "connection"
| "accept"
| "mcp-session-id"
| "mcp-protocol-version"
| "last-event-id"
) {
return Err(invalid(
"configuration",
"upstream header conflicts with HTTP framing",
));
}
let mut value = HeaderValue::try_from(value)
.map_err(|_| invalid("configuration", "invalid upstream header value"))?;
value.set_sensitive(true);
Ok((key, value))
}
fn scheme(value: &str, prefix: &str) -> String {
let bare = value
.get(..prefix.len())
.filter(|part| part.eq_ignore_ascii_case(prefix))
.and_then(|_| value.get(prefix.len()..))
.and_then(|suffix| suffix.strip_prefix(' '))
.unwrap_or(value);
format!("{prefix} {bare}")
}
fn basic(value: &str) -> String {
let normalized = scheme(value, "Basic");
let credentials = normalized.strip_prefix("Basic ").unwrap_or(value);
let encoded = if STANDARD.decode(credentials).is_ok() {
credentials.to_owned()
} else {
STANDARD.encode(credentials)
};
format!("Basic {encoded}")
}
async fn headers(
config: &McpServer,
secrets: &dyn SecretSource,
) -> Result<HashMap<HeaderName, HeaderValue>, ConnectError> {
let entries = try_join_all(
config
.static_headers
.iter()
.map(|(key, value)| async move { header(key, &resolve(value, secrets).await?) }),
)
.await?;
let auth = match (&config.authentication_token, config.auth_type) {
(Some(token), Some(kind)) => {
let token = resolve(token, secrets).await?;
if token.trim().is_empty() {
return Err(invalid("configuration", "upstream credential is empty"));
}
let (slot, value) = match kind {
McpAuth::ApiKey => ("x-api-key", token),
McpAuth::BearerToken => ("authorization", scheme(&token, "Bearer")),
McpAuth::Token => ("authorization", scheme(&token, "token")),
McpAuth::Basic => ("authorization", basic(&token)),
McpAuth::Authorization => ("authorization", token),
_ => return Err(invalid("configuration", "unsupported credential mode")),
};
Some(header(
config.upstream_token_header.as_deref().unwrap_or(slot),
&value,
)?)
}
_ => None,
};
Ok(entries.into_iter().chain(auth).collect())
}
impl ConfiguredGateway {
pub async fn connect<C: StreamableHttpClient>(
configs: &BTreeMap<String, McpServer>,
client: C,
secrets: &dyn SecretSource,
shutdown: CancellationToken,
) -> Result<Self, ConnectError> {
for (name, config) in configs {
validate(name, config)?;
}
let connected = try_join_all(configs.iter().map(|(name, config)| {
connect_one(
name,
config,
client.clone(),
secrets,
shutdown.child_token(),
)
}))
.await?;
let (servers, services, limits) = connected.into_iter().fold(
(Vec::new(), Vec::new(), BTreeMap::new()),
|(mut servers, mut services, mut limits), (server, running, limit)| {
limits.insert(server.info.server_id.clone(), limit);
servers.push(server);
services.push(running);
(servers, services, limits)
},
);
let registry = Registry::new(servers)?;
Ok(Self {
operations: Arc::new(
NativeGateway::new(Arc::new(registry), Duration::from_secs(30)).with_limits(limits),
),
services,
})
}
pub async fn close(self) {
futures_util::future::join_all(self.services.into_iter().map(|service| service.cancel()))
.await;
}
}
async fn connect_one<C: StreamableHttpClient>(
name: &str,
config: &McpServer,
client: C,
secrets: &dyn SecretSource,
shutdown: CancellationToken,
) -> Result<(Server, RunningService<RoleClient, ()>, Limits), ConnectError> {
let duration = timeout(name, config)?;
let running = tokio::select! {
result = tokio::time::timeout(duration, connect_service(name, config, client, secrets, shutdown.clone())) => result.map_err(|_| ConnectError::Timeout(name.into()))??,
() = shutdown.cancelled() => return Err(ConnectError::Cancelled),
};
let server = Server {
info: info(name, config),
peer: running.peer().clone(),
allowed_tools: config
.allowed_tools
.as_deref()
.filter(|tools| !tools.is_empty())
.map(Arc::from),
};
let limits = Limits::new(
duration,
config.max_concurrent_requests.and_then(NonZeroUsize::new),
);
Ok((server, running, limits))
}
async fn connect_service<C: StreamableHttpClient>(
name: &str,
config: &McpServer,
client: C,
secrets: &dyn SecretSource,
shutdown: CancellationToken,
) -> Result<RunningService<RoleClient, ()>, ConnectError> {
let initialized = match config.transport {
McpTransport::Http => {
let configured = config
.url
.as_ref()
.ok_or_else(|| invalid(name, "HTTP transport requires url"))?;
let url = resolve(configured, secrets).await?;
let parsed = url::Url::parse(&url).map_err(|_| invalid(name, "invalid HTTP URL"))?;
if !matches!(parsed.scheme(), "http" | "https")
|| parsed.host_str().is_none()
|| !parsed.username().is_empty()
|| parsed.password().is_some()
|| parsed.fragment().is_some()
{
return Err(invalid(
name,
"HTTP URL must use http or https without user credentials or fragments",
));
}
let transport = StreamableHttpClientTransport::with_client(
client,
StreamableHttpClientTransportConfig::with_uri(url)
.custom_headers(headers(config, secrets).await?),
);
().serve_with_ct(transport, shutdown).await
}
McpTransport::Stdio => {
let environment = try_join_all(config.env.iter().map(|(key, value)| async move {
Ok::<_, ConnectError>((key, resolve(value, secrets).await?))
}))
.await?;
let program = config
.command
.as_deref()
.ok_or_else(|| invalid(name, "stdio requires command"))?;
let mut command = tokio::process::Command::new(program);
command
.args(&config.args)
.envs(environment)
.kill_on_drop(true);
().serve_with_ct(TokioChildProcess::new(command)?, shutdown)
.await
}
McpTransport::Sse => return Err(invalid(name, "unsupported transport")),
};
initialized.map_err(|source| ConnectError::Initialize {
server: name.into(),
source: Box::new(source),
})
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
struct Secrets;
impl SecretSource for Secrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> futures_util::future::BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>>
{
Box::pin(async move { Ok((name == "KEY").then(|| SecretValue::new("resolved-token"))) })
}
}
#[rstest]
#[case::api_key("api_key", "token-value", "x-api-key", "token-value")]
#[case::bearer("bearer_token", "token-value", "authorization", "Bearer token-value")]
#[case::prefixed_bearer(
"bearer_token",
"bEaReR token-value",
"authorization",
"Bearer token-value"
)]
#[case::token("token", "token token-value", "authorization", "token token-value")]
#[case::basic(
"basic",
"user:password",
"authorization",
"Basic dXNlcjpwYXNzd29yZA=="
)]
#[case::prefixed_basic(
"basic",
"Basic dXNlcjpwYXNzd29yZA==",
"authorization",
"Basic dXNlcjpwYXNzd29yZA=="
)]
#[case::authorization(
"authorization",
"Custom token-value",
"authorization",
"Custom token-value"
)]
#[case::secret_reference("api_key", "os.environ/KEY", "x-api-key", "resolved-token")]
#[tokio::test]
async fn builds_configured_auth_header(
#[case] mode: &str,
#[case] token: &str,
#[case] slot: &str,
#[case] expected: &str,
) {
let config = litellm_config::Config::from_yaml(&format!(
"mcp_servers:\n test:\n auth_type: {mode}\n authentication_token: '{token}'\n"
))
.unwrap();
let built = headers(&config.mcp_servers["test"], &Secrets)
.await
.unwrap();
let key = HeaderName::try_from(slot).unwrap();
assert_eq!(built[&key], expected);
assert!(built[&key].is_sensitive());
}
#[rstest]
#[tokio::test]
async fn explicit_credential_slot_overrides_static_header() {
let config = litellm_config::Config::from_yaml("mcp_servers: {test: {auth_type: api_key, authentication_token: os.environ/KEY, upstream_token_header: x-custom, static_headers: {X-Custom: old, x-extra: extra}}}").unwrap();
let built = headers(&config.mcp_servers["test"], &Secrets)
.await
.unwrap();
assert_eq!(
built[&HeaderName::from_static("x-custom")],
"resolved-token"
);
assert_eq!(built[&HeaderName::from_static("x-extra")], "extra");
assert_eq!(built.len(), 2);
}
}

View file

@ -0,0 +1,98 @@
use axum::{
Json,
http::StatusCode,
response::{IntoResponse, Response},
};
use rmcp::{
model::{ErrorCode, ErrorData},
service::ServiceError,
};
use serde_json::json;
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("{0}")]
InvalidRequest(String),
#[error("User not allowed to access this MCP server or operation.")]
Forbidden,
#[error("{0}")]
Configuration(String),
#[error("MCP upstream request failed")]
Upstream(#[from] ServiceError),
#[error("MCP operation returned an unexpected result")]
UnexpectedResult,
#[error("MCP request cancelled")]
Cancelled,
}
impl Error {
pub fn into_mcp(self) -> ErrorData {
match self {
Self::Upstream(ServiceError::McpError(error)) => error,
Self::InvalidRequest(message) => ErrorData::invalid_params(message, None),
Self::Forbidden => ErrorData::new(ErrorCode(-32003), self.to_string(), None),
_ => ErrorData::internal_error(self.to_string(), None),
}
}
}
impl IntoResponse for Error {
fn into_response(self) -> Response {
let status = match &self {
Self::InvalidRequest(_) => StatusCode::BAD_REQUEST,
Self::Forbidden => StatusCode::FORBIDDEN,
Self::Upstream(_) => StatusCode::BAD_GATEWAY,
Self::Cancelled => StatusCode::REQUEST_TIMEOUT,
Self::Configuration(_) | Self::UnexpectedResult => StatusCode::INTERNAL_SERVER_ERROR,
};
(status, Json(json!({"detail": self.to_string()}))).into_response()
}
}
#[derive(thiserror::Error)]
pub enum ConnectError {
#[error("MCP server {server}: {message}")]
Configuration { server: String, message: String },
#[error("could not resolve MCP credentials")]
Secret(#[from] litellm_secrets::Error),
#[error("could not start MCP child process")]
Process(#[from] std::io::Error),
#[error("could not initialize MCP upstream {server}")]
Initialize {
server: String,
#[source]
source: Box<rmcp::service::ClientInitializeError>,
},
#[error("MCP upstream {0} initialization timed out")]
Timeout(String),
#[error("MCP upstream startup cancelled")]
Cancelled,
#[error(transparent)]
Registry(#[from] Error),
}
impl std::fmt::Debug for ConnectError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(self, formatter)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[rstest::rstest]
fn startup_diagnostics_do_not_expose_upstream_details() {
let error = ConnectError::Initialize {
server: "docs".into(),
source: Box::new(rmcp::service::ClientInitializeError::ConnectionClosed(
"https://example.test/mcp?token=private-secret".into(),
)),
};
assert_eq!(
format!("{error:?}"),
"could not initialize MCP upstream docs"
);
assert!(std::error::Error::source(&error).is_some());
}
}

View file

@ -0,0 +1,256 @@
use std::{sync::Arc, time::Duration};
use axum::{
Json, Router,
extract::{Query, Request, State},
http::{StatusCode, request::Parts},
response::{
IntoResponse, Response, Sse,
sse::{Event, KeepAlive},
},
routing::{get, post},
};
use futures_util::{Stream, StreamExt, sink, stream};
use moka::future::Cache;
use rmcp::{ServiceExt, model::*};
use serde::Deserialize;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use uuid::Uuid;
use crate::{Context, HttpConfig, McpServer, Operations, sessions::fingerprint};
#[derive(Clone)]
struct Session {
sender: mpsc::Sender<ClientJsonRpcMessage>,
owner: [u8; 32],
cancellation: CancellationToken,
}
#[derive(Clone)]
struct Legacy {
operations: Arc<dyn Operations>,
server: McpServer,
sessions: Cache<String, Session>,
shutdown: CancellationToken,
hosts: Arc<[String]>,
origins: Arc<[String]>,
}
pub(crate) fn router(
operations: Arc<dyn Operations>,
server: McpServer,
config: &HttpConfig,
) -> Router {
let state = Legacy {
operations,
server,
sessions: Cache::builder()
.max_capacity(10_000)
.time_to_idle(Duration::from_secs(300))
.eviction_listener(|_, session: Session, _| session.cancellation.cancel())
.build(),
shutdown: config.cancellation_token.clone(),
hosts: config.allowed_hosts.clone().into(),
origins: config.allowed_origins.clone().into(),
};
Router::new()
.route("/mcp/sse", get(connect))
.route("/mcp/sse/", get(connect))
.route("/mcp/sse/messages", post(message))
.route("/mcp/sse/messages/", post(message))
.with_state(state)
}
impl Legacy {
fn validate(&self, parts: &Parts) -> Result<(), StatusCode> {
let host = parts
.headers
.get("host")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<http::uri::Authority>().ok())
.ok_or(StatusCode::FORBIDDEN)?;
let allowed = self.hosts.iter().any(|allowed| {
host.as_str().eq_ignore_ascii_case(allowed)
|| host
.host()
.trim_matches(['[', ']'])
.eq_ignore_ascii_case(allowed)
});
if !allowed {
return Err(StatusCode::FORBIDDEN);
}
if let Some(origin) = parts.headers.get("origin") {
let origin = origin.to_str().map_err(|_| StatusCode::FORBIDDEN)?;
if !self
.origins
.iter()
.any(|allowed| same_origin(origin, allowed))
{
return Err(StatusCode::FORBIDDEN);
}
}
Ok(())
}
}
fn same_origin(value: &str, allowed: &str) -> bool {
if value == "null" {
return allowed == "null";
}
match (url::Url::parse(value), url::Url::parse(allowed)) {
(Ok(value), Ok(allowed)) => {
value.path() == "/"
&& value.query().is_none()
&& value.fragment().is_none()
&& value.username().is_empty()
&& value.password().is_none()
&& value.origin() == allowed.origin()
}
_ => false,
}
}
async fn connect(State(state): State<Legacy>, request: Request) -> Response {
let owner = fingerprint(&request);
let (parts, _) = request.into_parts();
if let Err(status) = state.validate(&parts) {
return status.into_response();
}
if let Err(error) = state
.operations
.authorize(Context {
parts,
server: None,
mcp: None,
})
.await
{
return error.into_response();
}
let connection = start_session(state, owner).await;
Sse::new(connection.events())
.keep_alive(KeepAlive::default())
.into_response()
}
struct Connection {
id: String,
output: mpsc::Receiver<ServerJsonRpcMessage>,
cancellation: CancellationToken,
}
async fn start_session(state: Legacy, owner: [u8; 32]) -> Connection {
let id = Uuid::new_v4().to_string();
let cancellation = state.shutdown.child_token();
let (input_sender, input_receiver) = mpsc::channel(16);
let (output_sender, output_receiver) = mpsc::channel(16);
state
.sessions
.insert(
id.clone(),
Session {
sender: input_sender,
owner,
cancellation: cancellation.clone(),
},
)
.await;
let connection = Connection {
id: id.clone(),
output: output_receiver,
cancellation: cancellation.clone(),
};
tokio::spawn(async move {
let sink = Box::pin(sink::unfold(output_sender, |sender, message| async move {
sender.send(message).await.map_err(std::io::Error::other)?;
Ok::<_, std::io::Error>(sender)
}));
let input = Box::pin(stream::unfold(input_receiver, |mut receiver| async move {
receiver.recv().await.map(|message| (message, receiver))
}));
if let Ok(service) = state
.server
.serve_with_ct((sink, input), cancellation)
.await
{
let _ = service.waiting().await;
}
state.sessions.invalidate(&id).await;
});
connection
}
impl Connection {
fn events(self) -> impl Stream<Item = Result<Event, std::io::Error>> {
let endpoint = Event::default()
.event("endpoint")
.data(format!("/mcp/sse/messages?session_id={}", self.id));
let first = stream::once(async move { Ok(endpoint) });
let messages = stream::unfold(
(self.output, self.cancellation.drop_guard()),
|(mut receiver, guard)| async move {
receiver
.recv()
.await
.map(|message| (message_event(message), (receiver, guard)))
},
);
first.chain(messages)
}
}
fn message_event(message: ServerJsonRpcMessage) -> Result<Event, std::io::Error> {
Event::default()
.event("message")
.json_data(message)
.map_err(std::io::Error::other)
}
#[derive(Deserialize)]
struct SessionQuery {
session_id: String,
}
async fn message(
State(state): State<Legacy>,
Query(query): Query<SessionQuery>,
request: Request,
) -> Response {
let owner = fingerprint(&request);
let (parts, body) = request.into_parts();
if let Err(status) = state.validate(&parts) {
return status.into_response();
}
let Some(session) = state.sessions.get(&query.session_id).await else {
return StatusCode::NOT_FOUND.into_response();
};
if session.owner != owner {
return StatusCode::FORBIDDEN.into_response();
}
let request = Request::from_parts(parts.clone(), body);
let Json(mut message) =
match <Json<ClientJsonRpcMessage> as axum::extract::FromRequest<()>>::from_request(
request,
&(),
)
.await
{
Ok(message) => message,
Err(error) => return error.into_response(),
};
match &mut message {
ClientJsonRpcMessage::Request(request) => {
request.request.extensions_mut().insert(parts);
}
ClientJsonRpcMessage::Notification(notification) => {
notification.notification.extensions_mut().insert(parts);
}
_ => (),
}
match session.sender.try_send(message) {
Ok(()) => (StatusCode::ACCEPTED, "Accepted").into_response(),
Err(mpsc::error::TrySendError::Full(_)) => StatusCode::TOO_MANY_REQUESTS.into_response(),
Err(mpsc::error::TrySendError::Closed(_)) => StatusCode::NOT_FOUND.into_response(),
}
}

View file

@ -0,0 +1,88 @@
mod configured;
mod error;
mod legacy;
mod native;
mod relay;
mod routes;
mod server;
mod sessions;
use std::{collections::BTreeMap, future::Future, pin::Pin, sync::Arc};
use http::request::Parts;
use rmcp::{RoleServer, model::*, service::RequestContext};
pub use configured::ConfiguredGateway;
pub use error::{ConnectError, Error};
pub use native::{Limits, NativeGateway, Registry, Server, ServerResolver};
pub use relay::RelayClient;
pub use rmcp;
pub use routes::{HttpConfig, router};
pub use server::McpServer;
pub use sessions::SessionOwner;
pub type GatewayFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
pub struct Context {
pub parts: Parts,
pub server: Option<String>,
pub mcp: Option<RequestContext<RoleServer>>,
}
#[derive(Clone, Debug)]
pub enum Operation {
ListTools(Option<PaginatedRequestParams>),
CallTool(CallToolRequestParams),
ListPrompts(Option<PaginatedRequestParams>),
GetPrompt(GetPromptRequestParams),
ListResources(Option<PaginatedRequestParams>),
ListResourceTemplates(Option<PaginatedRequestParams>),
ReadResource(ReadResourceRequestParams),
}
#[derive(serde::Serialize)]
pub struct RestTool {
#[serde(flatten)]
pub tool: Tool,
pub mcp_info: ServerInfo,
}
#[derive(serde::Serialize)]
pub struct ToolCatalog {
pub tools: Vec<RestTool>,
pub server_outcomes: BTreeMap<String, ServerOutcome>,
}
#[derive(serde::Serialize)]
#[serde(tag = "status", rename_all = "snake_case")]
pub enum ServerOutcome {
Ok { tool_count: usize },
Timeout,
Unreachable,
UpstreamError,
Internal,
}
#[derive(Clone, serde::Serialize)]
pub struct ServerInfo {
pub server_id: String,
pub server_name: String,
pub alias: Option<String>,
}
pub trait Operations: Send + Sync + 'static {
fn authorize(&self, context: Context) -> GatewayFuture<'_, ()>;
fn execute(&self, operation: Operation, context: Context) -> GatewayFuture<'_, ServerResult>;
fn rest_tools(&self, context: Context) -> GatewayFuture<'_, ToolCatalog>;
}
pub trait OperationAuthorizer: Send + Sync {
fn authorize<'a>(
&'a self,
server: &'a ServerInfo,
request: Option<&'a ClientRequest>,
) -> GatewayFuture<'a, ()>;
}
#[derive(Clone)]
pub struct Authorization(pub Arc<dyn OperationAuthorizer>);

View file

@ -0,0 +1,402 @@
mod catalog;
mod registry;
pub use registry::{Registry, Server, ServerResolver};
use std::{collections::BTreeMap, sync::Arc, time::Duration};
use tokio::sync::Semaphore;
use futures_util::future::join_all;
use rmcp::{model::*, service::PeerRequestOptions};
use crate::{Context, Error, GatewayFuture, Operation, Operations, ToolCatalog};
pub struct Limits {
timeout: Duration,
concurrency: Option<Semaphore>,
}
impl Limits {
pub fn new(timeout: Duration, concurrency: Option<std::num::NonZeroUsize>) -> Self {
Self {
timeout,
concurrency: concurrency.map(|limit| Semaphore::new(limit.get())),
}
}
}
pub struct NativeGateway {
resolver: Arc<dyn ServerResolver>,
timeout: Duration,
limits: BTreeMap<String, Limits>,
}
impl NativeGateway {
pub fn new(resolver: Arc<dyn ServerResolver>, timeout: Duration) -> Self {
Self {
resolver,
timeout,
limits: BTreeMap::new(),
}
}
pub fn with_limits(self, limits: BTreeMap<String, Limits>) -> Self {
Self { limits, ..self }
}
async fn send(
&self,
server: &Server,
mut request: ClientRequest,
context: &Context,
) -> Result<ServerResult, Error> {
if context
.mcp
.as_ref()
.is_some_and(|mcp| mcp.ct.is_cancelled())
{
return Err(Error::Cancelled);
}
if let Some(mcp) = &context.mcp {
request.get_meta_mut().extend(
mcp.meta
.iter()
.filter(|(key, _)| {
!matches!(
key.as_str(),
"progressToken"
| "io.modelcontextprotocol/protocolVersion"
| "io.modelcontextprotocol/clientInfo"
| "io.modelcontextprotocol/clientCapabilities"
)
})
.map(|(key, value)| (key.clone(), value.clone()))
.collect::<JsonObject>()
.into(),
);
}
if let Some(policy) = context.parts.extensions.get::<crate::Authorization>() {
policy.0.authorize(&server.info, Some(&request)).await?;
}
let limits = self.limits.get(&server.info.server_id);
let timeout = limits.map_or(self.timeout, |limits| limits.timeout);
let acquire = async {
match limits.and_then(|limits| limits.concurrency.as_ref()) {
Some(semaphore) => semaphore
.acquire()
.await
.map(Some)
.map_err(|_| Error::Cancelled),
None => Ok(None),
}
};
let cancelled = async {
match &context.mcp {
Some(mcp) => mcp.ct.cancelled().await,
None => std::future::pending().await,
}
};
let _permit = tokio::select! {
permit = tokio::time::timeout(timeout, acquire) => permit.map_err(|_| Error::Upstream(rmcp::service::ServiceError::Timeout { timeout }))??,
() = cancelled => return Err(Error::Cancelled),
};
let options = PeerRequestOptions::with_timeout(timeout);
let handle = server
.peer
.send_request_with_option(request, options)
.await?;
let request_id = handle.id.clone();
let cancelled = async {
match &context.mcp {
Some(mcp) => mcp.ct.cancelled().await,
None => std::future::pending().await,
}
};
tokio::select! {
response = handle.await_response() => Ok(response?),
() = cancelled => {
let notification = CancelledNotification::new(CancelledNotificationParam::new(Some(request_id), None));
let _ = server.peer.send_notification(notification.into()).await;
Err(Error::Cancelled)
}
}
}
async fn execute_native(
&self,
operation: Operation,
context: Context,
) -> Result<ServerResult, Error> {
let servers = self.servers(&context).await?;
match operation {
Operation::ListTools(params) => {
reject_cursor(&params)?;
let catalog = self.tools(&servers, &context, false).await?;
let tools = catalog
.tools
.into_iter()
.map(|entry| {
let prefix = entry
.mcp_info
.alias
.as_deref()
.unwrap_or(&entry.mcp_info.server_name)
.replace(' ', "_");
let mut tool = entry.tool;
tool.name = format!("{prefix}-{}", tool.name).into();
tool
})
.collect();
let mut result = ListToolsResult::with_all_items(tools);
result.meta = Some(MetaObject(JsonObject::from_iter([(
"litellm.ai/server_outcomes".into(),
serde_json::json!(catalog.server_outcomes),
)])));
Ok(result.into())
}
Operation::CallTool(request) => {
let catalog = self.tools(&servers, &context, false).await?;
let candidates: Vec<_> = catalog
.tools
.iter()
.filter(|entry| {
let prefix = entry
.mcp_info
.alias
.as_deref()
.unwrap_or(&entry.mcp_info.server_name)
.replace(' ', "_");
request.name == format!("{prefix}-{}", entry.tool.name)
|| request.name == entry.tool.name
})
.collect();
let [tool] = candidates.as_slice() else {
return Err(Error::Forbidden);
};
let server = servers
.iter()
.find(|server| server.info.server_id == tool.mcp_info.server_id)
.ok_or(Error::Forbidden)?;
let mut upstream = request;
upstream.name = tool.tool.name.clone();
match self
.send(server, CallToolRequest::new(upstream).into(), &context)
.await
{
Err(Error::Upstream(rmcp::service::ServiceError::McpError(error))) => {
Ok(CallToolResult::error(vec![ContentBlock::text(error.message)]).into())
}
result => result,
}
}
Operation::ListPrompts(params) => {
reject_cursor(&params)?;
let listings = join_all(servers.iter().map(|server| async {
if !server
.peer
.peer_info()
.is_some_and(|info| info.capabilities.prompts.is_some())
{
return Ok::<_, Error>(Vec::new());
}
let items = self
.pages(
server,
&context,
|params| {
ListPromptsRequest {
params,
..Default::default()
}
.into()
},
|result| match result {
ServerResult::ListPromptsResult(result) => {
Ok((result.prompts, result.next_cursor))
}
_ => Err(Error::UnexpectedResult),
},
)
.await?;
Ok(items
.into_iter()
.map(|mut item| {
item.name = format!("{}-{}", server.prefix(), item.name);
item
})
.collect::<Vec<_>>())
}))
.await;
Ok(ListPromptsResult::with_all_items(collect_listings(listings)?).into())
}
Operation::GetPrompt(request) => {
let (server, name) =
resolve_name(&servers, &request.name, context.server.is_some())?;
let mut upstream = request;
upstream.name = name;
self.send(server, GetPromptRequest::new(upstream).into(), &context)
.await
}
Operation::ListResources(params) => {
reject_cursor(&params)?;
let listings = join_all(servers.iter().map(|server| async {
if !server
.peer
.peer_info()
.is_some_and(|info| info.capabilities.resources.is_some())
{
return Ok::<_, Error>(Vec::new());
}
let items = self
.pages(
server,
&context,
|params| {
ListResourcesRequest {
params,
..Default::default()
}
.into()
},
|result| match result {
ServerResult::ListResourcesResult(result) => {
Ok((result.resources, result.next_cursor))
}
_ => Err(Error::UnexpectedResult),
},
)
.await?;
Ok(items
.into_iter()
.map(|mut item| {
item.name = format!("{}-{}", server.prefix(), item.name);
item
})
.collect::<Vec<_>>())
}))
.await;
Ok(ListResourcesResult::with_all_items(collect_listings(listings)?).into())
}
Operation::ListResourceTemplates(params) => {
reject_cursor(&params)?;
let listings = join_all(servers.iter().map(|server| async {
if !server
.peer
.peer_info()
.is_some_and(|info| info.capabilities.resources.is_some())
{
return Ok::<_, Error>(Vec::new());
}
let items = self
.pages(
server,
&context,
|params| {
ListResourceTemplatesRequest {
params,
..Default::default()
}
.into()
},
|result| match result {
ServerResult::ListResourceTemplatesResult(result) => {
Ok((result.resource_templates, result.next_cursor))
}
_ => Err(Error::UnexpectedResult),
},
)
.await?;
Ok(items
.into_iter()
.map(|mut item| {
item.name = format!("{}-{}", server.prefix(), item.name);
item
})
.collect::<Vec<_>>())
}))
.await;
Ok(ListResourceTemplatesResult::with_all_items(collect_listings(listings)?).into())
}
Operation::ReadResource(request) => {
if servers.is_empty() {
return Err(Error::Forbidden);
}
let [server] = servers.as_ref() else {
return Err(Error::InvalidRequest(
"read_resource requires exactly one allowed MCP server".into(),
));
};
self.send(server, ReadResourceRequest::new(request).into(), &context)
.await
}
}
}
}
impl Operations for NativeGateway {
fn authorize(&self, context: Context) -> GatewayFuture<'_, ()> {
Box::pin(async move {
let servers = self.servers(&context).await?;
if servers.is_empty() {
return Err(Error::Forbidden);
}
if let Some(policy) = context.parts.extensions.get::<crate::Authorization>() {
for server in servers.iter() {
policy.0.authorize(&server.info, None).await?;
}
}
Ok(())
})
}
fn execute(&self, operation: Operation, context: Context) -> GatewayFuture<'_, ServerResult> {
Box::pin(self.execute_native(operation, context))
}
fn rest_tools(&self, context: Context) -> GatewayFuture<'_, ToolCatalog> {
Box::pin(async move {
let servers = self.servers(&context).await?;
self.tools(&servers, &context, context.server.is_some())
.await
})
}
}
fn reject_cursor(params: &Option<PaginatedRequestParams>) -> Result<(), Error> {
match params.as_ref().and_then(|params| params.cursor.as_ref()) {
Some(_) => Err(Error::InvalidRequest(
"Gateway listings do not accept an upstream cursor".into(),
)),
None => Ok(()),
}
}
fn resolve_name<'a>(
servers: &'a [Server],
name: &str,
scoped: bool,
) -> Result<(&'a Server, String), Error> {
let matched = servers
.iter()
.filter_map(|server| {
name.strip_prefix(&format!("{}-", server.prefix()))
.map(|bare| (server, bare.to_owned()))
})
.max_by_key(|(server, _)| server.prefix().len());
match matched {
Some(found) => Ok(found),
None if scoped && servers.len() == 1 => Ok((&servers[0], name.to_owned())),
None => Err(Error::Forbidden),
}
}
fn collect_listings<T>(listings: Vec<Result<Vec<T>, Error>>) -> Result<Vec<T>, Error> {
let healthy = listings
.into_iter()
.map(|result| match result {
Err(Error::Cancelled) => Err(Error::Cancelled),
Err(_) => Ok(Vec::new()),
Ok(items) => Ok(items),
})
.collect::<Result<Vec<_>, _>>()?;
Ok(healthy.into_iter().flatten().collect())
}

View file

@ -0,0 +1,124 @@
use super::{NativeGateway, Server};
use crate::{Context, Error, RestTool, ServerOutcome, ToolCatalog};
use futures_util::{TryStreamExt, future::join_all, stream};
use rmcp::model::*;
use std::collections::BTreeSet;
impl NativeGateway {
pub(super) async fn pages<T, F, G>(
&self,
server: &Server,
context: &Context,
request: F,
unpack: G,
) -> Result<Vec<T>, Error>
where
T: Send,
F: Fn(Option<PaginatedRequestParams>) -> ClientRequest,
G: Fn(ServerResult) -> Result<(Vec<T>, Option<String>), Error>,
{
let pages: Vec<Vec<T>> = stream::try_unfold(Some((None, BTreeSet::new())), |state| async {
let Some((cursor, seen)) = state else {
return Ok(None);
};
let params =
cursor.map(|cursor| PaginatedRequestParams::default().with_cursor(Some(cursor)));
let result = self.send(server, request(params), context).await?;
let (items, next) = unpack(result)?;
let state = match next {
None => None,
Some(cursor) if seen.contains(&cursor) || seen.len() >= 1000 => {
return Err(Error::InvalidRequest(
"MCP upstream returned invalid pagination".into(),
));
}
Some(cursor) => Some((
Some(cursor.clone()),
seen.into_iter().chain([cursor]).collect(),
)),
};
Ok(Some((items, state)))
})
.try_collect()
.await?;
Ok(pages.into_iter().flatten().collect())
}
pub(super) async fn tools(
&self,
servers: &[Server],
context: &Context,
strict: bool,
) -> Result<ToolCatalog, Error> {
let catalogs = join_all(servers.iter().map(|server| async {
if !server
.peer
.peer_info()
.is_some_and(|info| info.capabilities.tools.is_some())
{
return Ok::<_, Error>(Vec::new());
}
let tools = self
.pages(
server,
context,
|params| {
ListToolsRequest {
params,
..Default::default()
}
.into()
},
|result| match result {
ServerResult::ListToolsResult(result) => {
Ok((result.tools, result.next_cursor))
}
_ => Err(Error::UnexpectedResult),
},
)
.await?;
Ok(tools
.into_iter()
.filter(|tool| server.allows(&tool.name))
.map(|tool| RestTool {
tool,
mcp_info: server.info.clone(),
})
.collect::<Vec<_>>())
}))
.await;
let outcomes = servers
.iter()
.zip(&catalogs)
.map(|(server, result)| {
let outcome = match result {
Ok(tools) => ServerOutcome::Ok {
tool_count: tools.len(),
},
Err(Error::Upstream(rmcp::service::ServiceError::Timeout { .. })) => {
ServerOutcome::Timeout
}
Err(Error::Upstream(
rmcp::service::ServiceError::TransportClosed
| rmcp::service::ServiceError::TransportSend(_),
)) => ServerOutcome::Unreachable,
Err(Error::Upstream(_)) => ServerOutcome::UpstreamError,
Err(_) => ServerOutcome::Internal,
};
(server.prefix(), outcome)
})
.collect();
let tools = catalogs
.into_iter()
.map(|result| match result {
Err(error) if strict || matches!(error, Error::Cancelled) => Err(error),
Err(_) => Ok(Vec::new()),
Ok(tools) => Ok(tools),
})
.collect::<Result<Vec<_>, _>>()?;
Ok(ToolCatalog {
tools: tools.into_iter().flatten().collect(),
server_outcomes: outcomes,
})
}
}

View file

@ -0,0 +1,130 @@
use super::NativeGateway;
use crate::{Context, Error, GatewayFuture, ServerInfo};
use rmcp::{Peer, RoleClient};
use std::sync::Arc;
#[derive(Clone)]
pub struct Server {
pub info: ServerInfo,
pub peer: Peer<RoleClient>,
pub allowed_tools: Option<Arc<[String]>>,
}
impl Server {
pub(super) fn prefix(&self) -> String {
self.info
.alias
.as_deref()
.unwrap_or(&self.info.server_name)
.replace(' ', "_")
}
pub(super) fn matches(&self, name: &str) -> bool {
self.info.server_id == name
|| self.info.server_name == name
|| self.info.alias.as_deref() == Some(name)
}
pub(super) fn allows(&self, name: &str) -> bool {
self.allowed_tools
.as_ref()
.is_none_or(|allowed| allowed.iter().any(|entry| entry == name))
}
}
pub trait ServerResolver: Send + Sync + 'static {
fn resolve<'a>(&'a self, context: &'a Context) -> GatewayFuture<'a, Arc<[Server]>>;
}
pub struct Registry(Arc<[Server]>);
impl Registry {
pub fn new(servers: impl Into<Arc<[Server]>>) -> Result<Self, Error> {
let servers = servers.into();
for (index, server) in servers.iter().enumerate() {
let identifiers = [&server.info.server_id, &server.info.server_name];
if identifiers.iter().any(|value| value.is_empty()) || server.prefix().is_empty() {
return Err(Error::Configuration(
"MCP server identifiers must not be empty".into(),
));
}
if servers[..index].iter().any(|previous| {
previous.prefix() == server.prefix()
|| identifiers.iter().any(|value| previous.matches(value))
|| server
.info
.alias
.as_deref()
.is_some_and(|alias| previous.matches(alias))
}) {
return Err(Error::Configuration(
"MCP server identifiers and prefixes must be unique".into(),
));
}
}
Ok(Self(servers))
}
}
impl ServerResolver for Registry {
fn resolve<'a>(&'a self, _: &'a Context) -> GatewayFuture<'a, Arc<[Server]>> {
Box::pin(async { Ok(self.0.clone()) })
}
}
impl NativeGateway {
pub(super) async fn servers(&self, context: &Context) -> Result<Arc<[Server]>, Error> {
let authorized = self.resolver.resolve(context).await?;
let header = context
.parts
.headers
.get("x-mcp-servers")
.map(|value| {
value
.to_str()
.map_err(|_| Error::InvalidRequest("Invalid x-mcp-servers header".into()))
})
.transpose()?;
let scope = context.server.as_deref().or(header);
let Some(scope) = scope else {
return Ok(authorized);
};
let names: Vec<_> = scope
.split(',')
.map(str::trim)
.filter(|name| !name.is_empty())
.collect();
if names.is_empty()
|| names
.iter()
.any(|name| !authorized.iter().any(|server| server.matches(name)))
{
return Err(Error::Forbidden);
}
let selected: Arc<[Server]> = authorized
.iter()
.filter(|server| names.iter().any(|name| server.matches(name)))
.cloned()
.collect();
if let Some(header) = header {
let requested: Vec<_> = header
.split(',')
.map(str::trim)
.filter(|name| !name.is_empty())
.collect();
if requested.is_empty()
|| requested
.iter()
.any(|name| !selected.iter().any(|server| server.matches(name)))
{
return Err(Error::Forbidden);
}
return Ok(selected
.iter()
.filter(|server| requested.iter().any(|name| server.matches(name)))
.cloned()
.collect());
}
Ok(selected)
}
}

View file

@ -0,0 +1,84 @@
use rmcp::{
ClientHandler, ErrorData, Peer, RoleClient, RoleServer,
model::*,
service::{NotificationContext, RequestContext},
};
use crate::Error;
#[derive(Clone)]
pub struct RelayClient {
downstream: Peer<RoleServer>,
info: ClientConfig,
progress_token: Option<ProgressToken>,
}
impl RelayClient {
pub fn for_request(context: &RequestContext<RoleServer>) -> Self {
Self {
downstream: context.peer.clone(),
info: ClientConfig::new(
context.client_capabilities().unwrap_or_default(),
Implementation::new("litellm-mcp-gateway", "1.0.0"),
),
progress_token: context.meta.get_progress_token(),
}
}
}
impl ClientHandler for RelayClient {
fn get_info(&self) -> ClientConfig {
self.info.clone()
}
#[allow(
deprecated,
reason = "preserve sampling for legacy MCP clients supported by the Python gateway"
)]
async fn create_message(
&self,
params: CreateMessageRequestParams,
_: RequestContext<RoleClient>,
) -> Result<CreateMessageResult, ErrorData> {
self.downstream
.create_message(params)
.await
.map_err(|error| Error::Upstream(error).into_mcp())
}
#[allow(
deprecated,
reason = "preserve roots for legacy MCP clients supported by the Python gateway"
)]
async fn list_roots(
&self,
_: RequestContext<RoleClient>,
) -> Result<ListRootsResult, ErrorData> {
self.downstream
.list_roots()
.await
.map_err(|error| Error::Upstream(error).into_mcp())
}
async fn create_elicitation(
&self,
params: ElicitRequestParams,
_: RequestContext<RoleClient>,
) -> Result<ElicitResult, ErrorData> {
self.downstream
.create_elicitation(params)
.await
.map_err(|error| Error::Upstream(error).into_mcp())
}
async fn on_progress(
&self,
mut params: ProgressNotificationParam,
_: NotificationContext<RoleClient>,
) {
if let Some(token) = &self.progress_token {
params.progress_token = token.clone();
let _ = self.downstream.notify_progress(params).await;
}
}
}

View file

@ -0,0 +1,157 @@
use std::sync::Arc;
use axum::{
Json, Router,
extract::{Path, Query, Request, State},
http::request::Parts,
middleware::{self, Next},
response::{IntoResponse, Response},
routing::{get, post},
};
use rmcp::{model::*, transport::streamable_http_server::StreamableHttpServerConfig};
use serde::{Deserialize, Serialize};
use tokio_util::sync::CancellationToken;
use crate::{Context, Error, McpServer, Operation, Operations, ToolCatalog, server::ServerScope};
pub struct HttpConfig {
pub allowed_hosts: Vec<String>,
pub allowed_origins: Vec<String>,
pub cancellation_token: CancellationToken,
pub server_info: ServerConfig,
}
impl Default for HttpConfig {
fn default() -> Self {
Self {
allowed_hosts: vec!["localhost".into(), "127.0.0.1".into(), "::1".into()],
allowed_origins: Vec::new(),
cancellation_token: CancellationToken::new(),
server_info: ServerConfig::new(
ServerCapabilities::builder()
.enable_tools()
.enable_prompts()
.enable_resources()
.build(),
)
.with_server_info(Implementation::new("litellm-mcp-server", "1.0.0")),
}
}
}
pub fn router(operations: Arc<dyn Operations>, config: HttpConfig) -> Router {
let server = McpServer::new(operations.clone()).with_info(config.server_info.clone());
let legacy = crate::legacy::router(operations.clone(), server.clone(), &config);
let transport_config = StreamableHttpServerConfig::default()
.with_legacy_session_mode(false)
.with_allowed_hosts(config.allowed_hosts)
.with_allowed_origins(config.allowed_origins)
.enforce_origin_validation()
.with_cancellation_token(config.cancellation_token);
let service = crate::sessions::transport(server, transport_config);
let scoped = Router::new()
.route_service("/mcp/{server}", service.clone())
.route_service("/mcp/{server}/mcp", service.clone())
.route_service("/{server}/mcp", service.clone())
.route_service("/{server}/mcp/", service.clone())
.route_layer(middleware::from_fn(scope_server));
Router::new()
.route_service("/mcp", service.clone())
.route_service("/mcp/", service)
.merge(scoped)
.route(
"/mcp/enabled",
get(|| async { Json(Enabled { enabled: true }) }),
)
.route("/mcp-rest/tools/list", get(rest_list))
.route("/mcp-rest/tools/call", post(rest_call))
.with_state(operations)
.merge(legacy)
}
async fn scope_server(Path(server): Path<String>, mut request: Request, next: Next) -> Response {
request.extensions_mut().insert(ServerScope(server));
next.run(request).await
}
#[derive(Serialize)]
struct Enabled {
enabled: bool,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct ListQuery {
server_id: Option<String>,
mcp_server_name: Option<String>,
}
#[derive(Serialize)]
struct ToolList {
#[serde(flatten)]
catalog: ToolCatalog,
error: Option<String>,
message: String,
}
async fn rest_list(
State(operations): State<Arc<dyn Operations>>,
Query(query): Query<ListQuery>,
parts: Parts,
) -> Result<impl IntoResponse, Error> {
let context = Context {
parts,
server: query.server_id.or(query.mcp_server_name),
mcp: None,
};
let catalog = operations.rest_tools(context).await?;
let failed = catalog.tools.is_empty()
&& catalog
.server_outcomes
.values()
.any(|outcome| !matches!(outcome, crate::ServerOutcome::Ok { .. }));
Ok(Json(ToolList {
catalog,
error: failed.then(|| "partial_failure".into()),
message: if failed {
"Failed to get tools from servers"
} else {
"Successfully retrieved tools"
}
.into(),
}))
}
#[derive(Deserialize)]
struct ToolCall {
server_id: Option<String>,
name: Option<String>,
arguments: Option<JsonObject>,
}
async fn rest_call(
State(operations): State<Arc<dyn Operations>>,
parts: Parts,
Json(call): Json<ToolCall>,
) -> Result<impl IntoResponse, Error> {
let server = call
.server_id
.filter(|id| !id.is_empty())
.ok_or_else(|| Error::InvalidRequest("server_id is required in request body".into()))?;
let name = call
.name
.filter(|name| !name.is_empty())
.ok_or_else(|| Error::InvalidRequest("name is required in request body".into()))?;
let params =
CallToolRequestParams::new(name).with_arguments(call.arguments.unwrap_or_default());
let context = Context {
parts,
server: Some(server),
mcp: None,
};
Ok(Json(
operations
.execute(Operation::CallTool(params), context)
.await?,
))
}

View file

@ -0,0 +1,165 @@
use std::sync::Arc;
use http::{Request, request::Parts};
use rmcp::{ErrorData, RoleServer, ServerHandler, model::*, service::RequestContext};
use crate::{Context, Error, Operation, Operations};
#[derive(Clone)]
pub struct McpServer {
operations: Arc<dyn Operations>,
info: ServerConfig,
}
impl McpServer {
pub fn new(operations: Arc<dyn Operations>) -> Self {
Self {
operations,
info: ServerConfig::new(
ServerCapabilities::builder()
.enable_tools()
.enable_prompts()
.enable_resources()
.build(),
)
.with_server_info(Implementation::new("litellm-mcp-server", "1.0.0")),
}
}
pub fn with_info(self, info: ServerConfig) -> Self {
Self { info, ..self }
}
async fn execute(
&self,
operation: Operation,
mcp: RequestContext<RoleServer>,
) -> Result<ServerResult, ErrorData> {
self.operations
.execute(operation, context(mcp))
.await
.map_err(Error::into_mcp)
}
}
#[derive(Clone)]
pub(crate) struct ServerScope(pub String);
macro_rules! listing {
($method:ident, $operation:ident, $result:ident) => {
async fn $method(
&self,
request: Option<PaginatedRequestParams>,
context: RequestContext<RoleServer>,
) -> Result<$result, ErrorData> {
match self
.execute(Operation::$operation(request), context)
.await?
{
ServerResult::$result(result) => Ok(result),
_ => Err(Error::UnexpectedResult.into_mcp()),
}
}
};
}
impl ServerHandler for McpServer {
fn get_info(&self) -> ServerConfig {
self.info.clone()
}
async fn initialize(
&self,
request: InitializeRequestParams,
mcp: RequestContext<RoleServer>,
) -> Result<InitializeResult, ErrorData> {
mcp.peer.set_peer_info(request.clone());
self.operations
.authorize(context(mcp))
.await
.map_err(Error::into_mcp)?;
self.negotiate_initialize(&request)
}
async fn discover(&self, mcp: RequestContext<RoleServer>) -> Result<DiscoverResult, ErrorData> {
self.operations
.authorize(context(mcp))
.await
.map_err(Error::into_mcp)?;
Ok(DiscoverResult::from_server_info(
self.supported_protocol_versions().into_owned(),
self.get_info(),
))
}
listing!(list_tools, ListTools, ListToolsResult);
listing!(list_prompts, ListPrompts, ListPromptsResult);
listing!(list_resources, ListResources, ListResourcesResult);
listing!(
list_resource_templates,
ListResourceTemplates,
ListResourceTemplatesResult
);
async fn call_tool(
&self,
request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
match self.execute(Operation::CallTool(request), context).await? {
ServerResult::CallToolResult(result) => Ok(result.into()),
ServerResult::InputRequiredResult(result) => {
Ok(CallToolResponse::InputRequired(result))
}
_ => Err(Error::UnexpectedResult.into_mcp()),
}
}
async fn get_prompt(
&self,
request: GetPromptRequestParams,
context: RequestContext<RoleServer>,
) -> Result<GetPromptResponse, ErrorData> {
match self.execute(Operation::GetPrompt(request), context).await? {
ServerResult::GetPromptResult(result) => Ok(result.into()),
ServerResult::InputRequiredResult(result) => {
Ok(GetPromptResponse::InputRequired(result))
}
_ => Err(Error::UnexpectedResult.into_mcp()),
}
}
async fn read_resource(
&self,
request: ReadResourceRequestParams,
context: RequestContext<RoleServer>,
) -> Result<ReadResourceResponse, ErrorData> {
match self
.execute(Operation::ReadResource(request), context)
.await?
{
ServerResult::ReadResourceResult(result) => Ok(result.into()),
ServerResult::InputRequiredResult(result) => {
Ok(ReadResourceResponse::InputRequired(result))
}
_ => Err(Error::UnexpectedResult.into_mcp()),
}
}
}
fn context(mcp: RequestContext<RoleServer>) -> Context {
let parts = mcp
.extensions
.get::<Parts>()
.cloned()
.unwrap_or_else(|| Request::new(()).into_parts().0);
let server = parts
.extensions
.get::<ServerScope>()
.map(|scope| scope.0.clone());
Context {
parts,
server,
mcp: Some(mcp),
}
}

View file

@ -0,0 +1,157 @@
use std::{sync::Arc, time::Duration};
use axum::{
Router,
body::{Body, to_bytes},
extract::{Request, State},
http::{Method, StatusCode},
response::IntoResponse,
routing::any,
};
use moka::future::Cache;
use rmcp::{
model::{ClientJsonRpcMessage, ClientRequest},
transport::streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService,
session::{SessionManager, local::LocalSessionManager},
},
};
use sha2::{Digest, Sha256};
use tower::ServiceExt;
use crate::{McpServer, server::ServerScope};
#[derive(Clone)]
pub struct SessionOwner(pub String);
#[derive(Clone)]
struct Transport {
stateful: Router,
stateless: Router,
owners: Cache<String, [u8; 32]>,
}
pub(crate) fn transport(server: McpServer, config: StreamableHttpServerConfig) -> Router {
let manager = Arc::new(LocalSessionManager::default());
let cleanup = manager.clone();
let owners = Cache::builder()
.max_capacity(10_000)
.time_to_idle(Duration::from_secs(300))
.async_eviction_listener(move |id: Arc<String>, _, _| {
let manager = cleanup.clone();
Box::pin(async move {
let _ = manager.close_session(&id.as_str().into()).await;
})
})
.build();
let stateful_server = server.clone();
let stateful = StreamableHttpService::new(
move || Ok(stateful_server.clone()),
manager,
config.clone().with_legacy_session_mode(true),
);
let stateless = StreamableHttpService::new(
move || Ok(server.clone()),
Arc::new(LocalSessionManager::default()),
config.with_legacy_session_mode(false),
);
Router::new().fallback(any(dispatch)).with_state(Transport {
stateful: Router::new().fallback_service(stateful),
stateless: Router::new().fallback_service(stateless),
owners,
})
}
async fn dispatch(
State(transport): State<Transport>,
request: Request,
) -> Result<impl IntoResponse, StatusCode> {
let owner = fingerprint(&request);
let session = request
.headers()
.get("mcp-session-id")
.and_then(|value| value.to_str().ok())
.map(str::to_owned);
if let Some(session) = &session {
match transport.owners.get(session).await {
Some(expected) if expected == owner => (),
Some(_) => return Err(StatusCode::FORBIDDEN),
None => return Err(StatusCode::NOT_FOUND),
}
}
let method = request.method().clone();
let (request, initialize) = if method == Method::POST && session.is_none() {
let (parts, body) = request.into_parts();
let bytes = to_bytes(body, 4 * 1024 * 1024)
.await
.map_err(|_| StatusCode::PAYLOAD_TOO_LARGE)?;
let initialize = matches!(serde_json::from_slice::<ClientJsonRpcMessage>(&bytes),
Ok(ClientJsonRpcMessage::Request(request)) if matches!(request.request, ClientRequest::InitializeRequest(_)));
(Request::from_parts(parts, Body::from(bytes)), initialize)
} else {
(request, false)
};
let service = if initialize || session.is_some() {
transport.stateful
} else {
transport.stateless
};
let response = match service.oneshot(request).await {
Ok(response) => response,
Err(never) => match never {},
};
if initialize
&& response.status().is_success()
&& let Some(id) = response
.headers()
.get("mcp-session-id")
.and_then(|value| value.to_str().ok())
{
transport.owners.insert(id.to_owned(), owner).await;
}
if method == Method::DELETE
&& response.status().is_success()
&& let Some(session) = session
{
transport.owners.invalidate(&session).await;
}
Ok(response)
}
pub(crate) fn fingerprint(request: &Request) -> [u8; 32] {
let scope = request
.extensions()
.get::<ServerScope>()
.map(|scope| scope.0.as_str())
.unwrap_or_default();
let mut digest = Sha256::new();
match request.extensions().get::<SessionOwner>() {
Some(owner) => {
digest.update(b"principal");
digest.update(owner.0.len().to_be_bytes());
digest.update(owner.0.as_bytes());
}
None => {
digest.update(b"credentials");
for name in ["authorization", "x-litellm-api-key"] {
let value = request
.headers()
.get(name)
.map(|value| value.as_bytes())
.unwrap_or_default();
digest.update(value.len().to_be_bytes());
digest.update(value);
}
}
}
digest.update(scope.len().to_be_bytes());
digest.update(scope.as_bytes());
digest.update(
request
.headers()
.get("x-mcp-servers")
.map(|value| value.as_bytes())
.unwrap_or_default(),
);
digest.finalize().into()
}

View file

@ -0,0 +1,170 @@
use litellm_gateway_mcp::{Limits, McpServer, NativeGateway, Registry, Server, ServerInfo};
use rmcp::{
ErrorData, RoleServer, ServerHandler, ServiceExt,
model::*,
service::{PeerRequestOptions, RequestContext},
};
use rstest::rstest;
use std::{collections::BTreeMap, num::NonZeroUsize, sync::Arc, time::Duration};
use tokio::sync::Notify;
struct Slow {
entered: Arc<Notify>,
cancelled: Arc<Notify>,
}
struct Finished(Arc<Notify>);
impl Drop for Finished {
fn drop(&mut self) {
self.0.notify_one();
}
}
impl ServerHandler for Slow {
fn get_info(&self) -> ServerConfig {
ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
}
async fn list_tools(
&self,
_: Option<PaginatedRequestParams>,
_: RequestContext<RoleServer>,
) -> Result<ListToolsResult, ErrorData> {
Ok(ListToolsResult::with_all_items(vec![Tool::new(
"wait",
"Wait",
Arc::new(JsonObject::new()),
)]))
}
async fn call_tool(
&self,
_: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
let _finished = Finished(self.cancelled.clone());
self.entered.notify_one();
context.ct.cancelled().await;
Err(ErrorData::internal_error("cancelled", None))
}
}
#[rstest]
#[case::client_cancel(false, false)]
#[case::upstream_timeout(true, false)]
#[case::configured_timeout(true, true)]
#[case::configured_concurrency(false, true)]
#[tokio::test]
async fn cancellation_reaches_upstream(#[case] timeout: bool, #[case] configured: bool) {
let entered = Arc::new(Notify::new());
let cancelled = Arc::new(Notify::new());
let upstream_server = Slow {
entered: entered.clone(),
cancelled: cancelled.clone(),
};
let (server, client) = tokio::io::duplex(65536);
let upstream_task = tokio::spawn(async move {
upstream_server
.serve(server)
.await
.unwrap()
.waiting()
.await
.unwrap()
});
let upstream = ().serve(client).await.unwrap();
let registry = Registry::new(vec![Server {
info: ServerInfo {
server_id: "slow".into(),
server_name: "slow".into(),
alias: None,
},
peer: upstream.peer().clone(),
allowed_tools: None,
}])
.unwrap();
let duration = if timeout {
Duration::from_millis(100)
} else {
Duration::from_secs(5)
};
let native = NativeGateway::new(
Arc::new(registry),
if configured {
Duration::from_secs(5)
} else {
duration
},
);
let native = if configured {
native.with_limits(BTreeMap::from([(
"slow".into(),
Limits::new(duration, NonZeroUsize::new(1)),
)]))
} else {
native
};
let gateway = McpServer::new(Arc::new(native));
let (server, client) = tokio::io::duplex(65536);
let gateway_task = tokio::spawn(async move {
gateway
.serve(server)
.await
.unwrap()
.waiting()
.await
.unwrap()
});
let client = ().serve(client).await.unwrap();
let handle = client
.send_request_with_option(
CallToolRequest::new(CallToolRequestParams::new("slow-wait")).into(),
PeerRequestOptions::with_timeout(Duration::from_secs(2)),
)
.await
.unwrap();
tokio::time::timeout(Duration::from_secs(2), entered.notified())
.await
.unwrap();
let queued = if configured && !timeout {
let second = client
.send_request_with_option(
CallToolRequest::new(CallToolRequestParams::new("slow-wait")).into(),
PeerRequestOptions::with_timeout(Duration::from_secs(5)),
)
.await
.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(100), entered.notified())
.await
.is_err()
);
Some(second)
} else {
None
};
if timeout {
assert!(matches!(
handle.await_response().await,
Err(rmcp::ServiceError::McpError(_))
));
} else {
handle.cancel(None).await.unwrap();
}
tokio::time::timeout(Duration::from_secs(2), cancelled.notified())
.await
.unwrap();
if let Some(queued) = queued {
tokio::time::timeout(Duration::from_secs(2), entered.notified())
.await
.unwrap();
queued.cancel(None).await.unwrap();
tokio::time::timeout(Duration::from_secs(2), cancelled.notified())
.await
.unwrap();
}
client.cancel().await.unwrap();
upstream.cancel().await.unwrap();
gateway_task.await.unwrap();
upstream_task.await.unwrap();
}

View file

@ -0,0 +1,467 @@
use super::support;
use super::support::{Harness, harness};
use axum::{
body::{Body, to_bytes},
http::{Request, StatusCode},
};
use litellm_gateway_mcp::{
Context, Error, GatewayFuture, HttpConfig, NativeGateway, Registry, Server, ServerResolver,
router,
};
use rstest::rstest;
use serde_json::{Value, json};
use std::{sync::Arc, time::Duration};
use tower::ServiceExt;
#[rstest]
#[case::aggregate("/mcp", 4)]
#[case::trailing_slash("/mcp/", 4)]
#[case::scope_suffix("/alpha/mcp", 2)]
#[case::scope_prefix("/mcp/alpha", 2)]
#[case::server_id("/id-alpha/mcp", 2)]
#[tokio::test]
async fn streamable_http_lists_scoped_tools(
#[future] harness: Harness,
#[case] path: &str,
#[case] count: usize,
) {
let harness = harness.await;
let response = harness
.app
.oneshot(
Request::post(path)
.header("host", "localhost")
.header("content-type", "application/json")
.header("accept", "application/json, text/event-stream")
.header("mcp-protocol-version", "2025-11-25")
.body(Body::from(
json!({"jsonrpc":"2.0","id":42,"method":"tools/list"}).to_string(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = to_bytes(response.into_body(), 65536).await.unwrap();
let wire = String::from_utf8(body.to_vec()).unwrap();
let data = wire
.lines()
.find_map(|line| line.strip_prefix("data: "))
.unwrap();
let result: Value = serde_json::from_str(data).unwrap();
assert_eq!(result["id"], 42);
assert_eq!(result["result"]["tools"].as_array().unwrap().len(), count);
}
#[rstest]
#[tokio::test]
async fn rest_lists_bare_names_and_calls_selected_server(#[future] harness: Harness) {
let harness = harness.await;
let listing = harness
.app
.clone()
.oneshot(
Request::get("/mcp-rest/tools/list?server_id=id-alpha")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(listing.status(), StatusCode::OK);
let listing: Value =
serde_json::from_slice(&to_bytes(listing.into_body(), 65536).await.unwrap()).unwrap();
assert_eq!(listing["tools"][0]["name"], "echo-value");
assert_eq!(listing["tools"][0]["mcp_info"]["server_id"], "id-alpha");
assert_eq!(listing["tools"].as_array().unwrap().len(), 2);
let response = harness
.app
.oneshot(
Request::post("/mcp-rest/tools/call")
.header("content-type", "application/json")
.body(Body::from(
json!({"server_id":"id-alpha", "name":"echo-value", "arguments":{"foo":"bar"}})
.to_string(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let response: Value =
serde_json::from_slice(&to_bytes(response.into_body(), 65536).await.unwrap()).unwrap();
let forwarded: Value =
serde_json::from_str(response["content"][0]["text"].as_str().unwrap()).unwrap();
assert_eq!(forwarded["arguments"], json!({"foo":"bar"}));
}
#[rstest]
#[case::foreign_origin("localhost", Some("https://attacker.example"))]
#[case::dns_rebinding("attacker.example", None)]
#[tokio::test]
async fn rejects_untrusted_browser_origins_and_hosts(
#[future] harness: Harness,
#[case] host: &str,
#[case] origin: Option<&str>,
) {
let harness = harness.await;
let request = Request::post("/mcp")
.header("host", host)
.header("content-type", "application/json")
.header("accept", "application/json, text/event-stream");
let request = match origin {
Some(origin) => request.header("origin", origin),
None => request,
};
let response = harness
.app
.oneshot(
request
.body(Body::from(
json!({"jsonrpc":"2.0","id":1,"method":"tools/list"}).to_string(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::FORBIDDEN);
}
struct PerRequestResolver(Arc<[Server]>);
impl ServerResolver for PerRequestResolver {
fn resolve<'a>(&'a self, context: &'a Context) -> GatewayFuture<'a, Arc<[Server]>> {
Box::pin(async move {
match context
.parts
.headers
.get("authorization")
.and_then(|value| value.to_str().ok())
{
Some("Bearer allowed") => Ok(self.0.clone()),
_ => Err(Error::Forbidden),
}
})
}
}
#[rstest]
#[case::mcp("/mcp", "POST")]
#[case::rest("/mcp-rest/tools/list", "GET")]
#[tokio::test]
async fn resolver_receives_each_requests_identity(
#[future] harness: Harness,
#[case] path: &str,
#[case] method: &str,
) {
let harness = harness.await;
let server = Server {
info: litellm_gateway_mcp::ServerInfo {
server_id: "id-alpha".into(),
server_name: "alpha".into(),
alias: None,
},
peer: harness.connections[0].peer().clone(),
allowed_tools: None,
};
let registry = Registry::new(vec![server]).unwrap();
let servers = registry.resolve(&support::context(None)).await.unwrap();
let operations = Arc::new(NativeGateway::new(
Arc::new(PerRequestResolver(servers)),
Duration::from_secs(2),
));
let app = router(operations, HttpConfig::default());
for token in ["allowed", "denied", "allowed"] {
let response = app
.clone()
.oneshot(
Request::builder()
.method(method)
.uri(path)
.header("host", "localhost")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.header("accept", "application/json, text/event-stream")
.header("mcp-protocol-version", "2025-11-25")
.body(Body::from(
json!({"jsonrpc":"2.0","id":1,"method":"tools/list"}).to_string(),
))
.unwrap(),
)
.await
.unwrap();
let status = response.status();
let body = String::from_utf8(
to_bytes(response.into_body(), 65536)
.await
.unwrap()
.to_vec(),
)
.unwrap();
if token == "denied" {
if method == "GET" {
assert_eq!(status, StatusCode::FORBIDDEN);
}
assert!(body.contains("not allowed"), "{body}");
} else {
assert_eq!(status, StatusCode::OK);
assert!(body.contains("echo-value"), "{body}");
}
}
}
fn rpc_request(path: &str, token: &str, session: Option<&str>, body: Value) -> Request<Body> {
let request = Request::post(path)
.header("host", "localhost")
.header("authorization", format!("Bearer {token}"))
.header("content-type", "application/json")
.header("accept", "application/json, text/event-stream")
.header("mcp-protocol-version", "2025-11-25");
let request = match session {
Some(session) => request.header("mcp-session-id", session),
None => request,
};
request.body(Body::from(body.to_string())).unwrap()
}
#[rstest]
#[case::different_owner("other", "/alpha/mcp", false)]
#[case::forged_alternate_key("other", "/alpha/mcp", true)]
#[case::different_scope("owner", "/alpha-beta/mcp", false)]
#[tokio::test]
async fn binds_sessions_to_owner_and_scope(
#[future] harness: Harness,
#[case] token: &str,
#[case] path: &str,
#[case] forged: bool,
) {
let harness = harness.await;
let response = harness.app.clone().oneshot(rpc_request("/alpha/mcp", "owner", None, json!({
"jsonrpc":"2.0", "id":1, "method":"initialize", "params":{
"protocolVersion":"2025-11-25", "capabilities":{}, "clientInfo":{"name":"fixture","version":"1"}
}
}))).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let session = response.headers()["mcp-session-id"]
.to_str()
.unwrap()
.to_owned();
let initialized = harness
.app
.clone()
.oneshot(rpc_request(
"/alpha/mcp",
"owner",
Some(&session),
json!({"jsonrpc":"2.0","method":"notifications/initialized"}),
))
.await
.unwrap();
assert_eq!(initialized.status(), StatusCode::ACCEPTED);
let mut denied_request = rpc_request(
path,
token,
Some(&session),
json!({"jsonrpc":"2.0","id":2,"method":"tools/list"}),
);
if forged {
denied_request
.headers_mut()
.insert("x-litellm-api-key", "Bearer owner".parse().unwrap());
}
let denied = harness.app.clone().oneshot(denied_request).await.unwrap();
assert_eq!(denied.status(), StatusCode::FORBIDDEN);
let allowed = harness
.app
.clone()
.oneshot(rpc_request(
"/alpha/mcp",
"owner",
Some(&session),
json!({"jsonrpc":"2.0","id":3,"method":"tools/list"}),
))
.await
.unwrap();
assert_eq!(allowed.status(), StatusCode::OK);
let body =
String::from_utf8(to_bytes(allowed.into_body(), 65536).await.unwrap().to_vec()).unwrap();
assert!(body.contains("alpha-echo-value"), "{body}");
let delete = Request::delete("/alpha/mcp")
.header("host", "localhost")
.header("authorization", "Bearer owner")
.header("mcp-session-id", &session)
.body(Body::empty())
.unwrap();
assert!(
harness
.app
.clone()
.oneshot(delete)
.await
.unwrap()
.status()
.is_success()
);
let expired = harness
.app
.oneshot(rpc_request(
"/alpha/mcp",
"owner",
Some(&session),
json!({"jsonrpc":"2.0","id":4,"method":"tools/list"}),
))
.await
.unwrap();
assert_eq!(expired.status(), StatusCode::NOT_FOUND);
}
#[rstest]
#[case::narrowed("/mcp", "alpha", false)]
#[case::multiple("/mcp", "alpha,alpha-beta", false)]
#[case::path_broadening("/alpha/mcp", "alpha-beta", true)]
#[case::unknown("/mcp", "unknown", true)]
#[tokio::test]
async fn header_scope_cannot_broaden_path_scope(
#[future] harness: Harness,
#[case] path: &str,
#[case] scope: &str,
#[case] denied: bool,
) {
let harness = harness.await;
let mut request = rpc_request(
path,
"owner",
None,
json!({"jsonrpc":"2.0","id":1,"method":"tools/list"}),
);
request
.headers_mut()
.insert("x-mcp-servers", scope.parse().unwrap());
let response = harness.app.oneshot(request).await.unwrap();
let body = String::from_utf8(
to_bytes(response.into_body(), 65536)
.await
.unwrap()
.to_vec(),
)
.unwrap();
assert_eq!(body.contains("not allowed"), denied, "{body}");
if !denied {
assert!(body.contains("alpha-echo-value"));
}
}
#[rstest]
#[case::missing_server(json!({"name":"echo-value"}))]
#[case::missing_name(json!({"server_id":"alpha"}))]
#[tokio::test]
async fn rest_requires_server_and_tool(#[future] harness: Harness, #[case] body: Value) {
let harness = harness.await;
let response = harness
.app
.oneshot(
Request::post("/mcp-rest/tools/call")
.header("content-type", "application/json")
.body(Body::from(body.to_string()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
}
#[rstest]
#[tokio::test]
async fn legacy_sse_initializes_calls_tools_and_rejects_other_owners(#[future] harness: Harness) {
use futures_util::StreamExt;
let harness = harness.await;
let response = harness
.app
.clone()
.oneshot(
Request::get("/mcp/sse")
.header("host", "localhost")
.header("authorization", "Bearer owner")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let mut events = sse_stream::SseStream::new(response.into_body());
let endpoint = events.next().await.unwrap().unwrap();
assert_eq!(endpoint.event.as_deref(), Some("endpoint"));
let path = endpoint.data.unwrap();
let response = harness
.app
.clone()
.oneshot(rpc_request(
&path,
"other",
None,
json!({"jsonrpc":"2.0","id":1,"method":"tools/list"}),
))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::FORBIDDEN);
let response = harness.app.clone().oneshot(rpc_request(&path,"owner",None,json!({
"jsonrpc":"2.0", "id":1, "method":"initialize", "params":{
"protocolVersion":"2025-11-25", "capabilities":{}, "clientInfo":{"name":"fixture","version":"1"}
}
}))).await.unwrap();
assert_eq!(response.status(), StatusCode::ACCEPTED);
let initialized = events.next().await.unwrap().unwrap();
assert_eq!(initialized.event.as_deref(), Some("message"));
let initialized: Value = serde_json::from_str(&initialized.data.unwrap()).unwrap();
assert_eq!(
initialized["result"]["serverInfo"]["name"],
"litellm-mcp-server"
);
let response = harness
.app
.clone()
.oneshot(rpc_request(
&path,
"owner",
None,
json!({"jsonrpc":"2.0","method":"notifications/initialized"}),
))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::ACCEPTED);
let response = harness.app.clone().oneshot(rpc_request(&path,"owner",None,json!({"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"alpha-echo-value","arguments":{"source":"sse"}}}))).await.unwrap();
assert_eq!(response.status(), StatusCode::ACCEPTED);
let called = tokio::time::timeout(Duration::from_secs(2), events.next())
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(called.event.as_deref(), Some("message"));
let called: Value = serde_json::from_str(&called.data.unwrap()).unwrap();
let forwarded: Value =
serde_json::from_str(called["result"]["content"][0]["text"].as_str().unwrap()).unwrap();
assert_eq!(forwarded["arguments"], json!({"source":"sse"}));
drop(events);
tokio::time::timeout(Duration::from_secs(2), async {
loop {
let response = harness
.app
.clone()
.oneshot(rpc_request(
&path,
"owner",
None,
json!({"jsonrpc":"2.0","method":"notifications/initialized"}),
))
.await
.unwrap();
if response.status() == StatusCode::NOT_FOUND {
break;
}
assert_eq!(response.status(), StatusCode::ACCEPTED);
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("disconnected SSE session should reject subsequent messages");
}

View file

@ -0,0 +1,5 @@
mod cancellation;
mod http;
mod native;
mod relay;
mod support;

View file

@ -0,0 +1,353 @@
use std::sync::atomic::Ordering;
use super::support::{Harness, context, harness};
use litellm_gateway_mcp::{Error, McpServer, Operation, Operations};
use rmcp::{ServiceExt, model::*};
use rstest::rstest;
use serde_json::{Value, json};
#[rstest]
#[tokio::test]
async fn sdk_client_discovers_and_calls_prefixed_tools(#[future] harness: Harness) {
let harness = harness.await;
let (server, client) = tokio::io::duplex(65536);
let gateway = McpServer::new(harness.gateway.clone());
let server = tokio::spawn(async move {
gateway
.serve(server)
.await
.unwrap()
.waiting()
.await
.unwrap()
});
let client = ().serve(client).await.unwrap();
let tools = client.list_all_tools().await.unwrap();
let names: Vec<_> = tools.iter().map(|tool| tool.name.as_ref()).collect();
assert_eq!(
names,
[
"alpha-echo-value",
"alpha-fail",
"alpha-beta-echo-value",
"alpha-beta-fail"
]
);
let arguments = serde_json::from_value(json!({"value": [1, true, "hello"]})).unwrap();
let mut request = CallToolRequestParams::new("alpha-beta-echo-value").with_arguments(arguments);
request.meta = Some(serde_json::from_value(json!({"traceparent":"00-0123456789abcdef0123456789abcdef-0123456789abcdef-01", "custom":"preserved"})).unwrap());
let result = client.call_tool(request).await.unwrap();
let text = result.content[0].as_text().unwrap();
let forwarded: Value = serde_json::from_str(&text.text).unwrap();
assert_eq!(forwarded["name"], "echo-value");
assert_eq!(forwarded["arguments"], json!({"value": [1, true, "hello"]}));
assert_eq!(forwarded["meta"]["custom"], "preserved");
assert_eq!(harness.calls.load(Ordering::SeqCst), 1);
client.cancel().await.unwrap();
server.await.unwrap();
}
#[rstest]
#[case::hidden("alpha-hidden", None)]
#[case::unknown("unknown-echo-value", None)]
#[case::bare_ambiguous("echo-value", None)]
#[case::cross_scope("alpha-beta-echo-value", Some("alpha"))]
#[case::unknown_scope("echo-value", Some("missing"))]
#[tokio::test]
async fn denies_calls_outside_visible_catalog(
#[future] harness: Harness,
#[case] name: &str,
#[case] scope: Option<&str>,
) {
let harness = harness.await;
let result = harness
.gateway
.execute(
Operation::CallTool(CallToolRequestParams::new(name.to_owned())),
context(scope),
)
.await;
assert!(matches!(result, Err(Error::Forbidden)));
assert_eq!(harness.calls.load(Ordering::SeqCst), 0);
}
#[rstest]
#[tokio::test]
async fn preserves_tool_failures_as_tool_results(#[future] harness: Harness) {
let harness = harness.await;
let result = harness
.gateway
.execute(
Operation::CallTool(CallToolRequestParams::new("alpha-fail")),
context(None),
)
.await
.unwrap();
let ServerResult::CallToolResult(result) = result else {
panic!("expected tool result")
};
assert_eq!(result.is_error, Some(true));
assert_eq!(
result.content[0].as_text().unwrap().text,
"fixture tool failed"
);
}
#[rstest]
#[tokio::test]
async fn routes_prompts_and_resources_without_rewriting_uris(#[future] harness: Harness) {
let harness = harness.await;
let prompts = harness
.gateway
.execute(Operation::ListPrompts(None), context(Some("alpha-beta")))
.await
.unwrap();
let ServerResult::ListPromptsResult(prompts) = prompts else {
panic!("expected prompts")
};
assert_eq!(prompts.prompts[0].name, "alpha-beta-review-code");
let prompt = harness
.gateway
.execute(
Operation::GetPrompt(GetPromptRequestParams::new("alpha-beta-review-code")),
context(None),
)
.await
.unwrap();
let ServerResult::GetPromptResult(prompt) = prompt else {
panic!("expected prompt")
};
assert_eq!(
serde_json::to_value(prompt).unwrap()["messages"][0]["content"]["text"],
"review-code"
);
let resources = harness
.gateway
.execute(Operation::ListResources(None), context(Some("alpha")))
.await
.unwrap();
let ServerResult::ListResourcesResult(resources) = resources else {
panic!("expected resources")
};
assert_eq!(resources.resources[0].name, "alpha-document");
assert_eq!(resources.resources[0].uri, "fixture://document");
let templates = harness
.gateway
.execute(
Operation::ListResourceTemplates(None),
context(Some("alpha")),
)
.await
.unwrap();
let ServerResult::ListResourceTemplatesResult(templates) = templates else {
panic!("expected templates")
};
assert_eq!(templates.resource_templates[0].name, "alpha-files");
assert_eq!(
templates.resource_templates[0].uri_template,
"fixture://{file}"
);
let read = Operation::ReadResource(ReadResourceRequestParams::new("fixture://document"));
assert!(matches!(
harness.gateway.execute(read.clone(), context(None)).await,
Err(Error::InvalidRequest(_))
));
let result = harness
.gateway
.execute(read, context(Some("alpha")))
.await
.unwrap();
assert_eq!(
serde_json::to_value(result).unwrap()["contents"][0]["text"],
"fixture contents"
);
}
#[rstest]
#[tokio::test]
async fn retains_healthy_catalogs_and_reports_failed_servers(#[future] harness: Harness) {
let harness = harness.await;
let [first, second] = <[_; 2]>::try_from(harness.connections).ok().unwrap();
second.cancel().await.unwrap();
let result = harness
.gateway
.execute(Operation::ListTools(None), context(None))
.await
.unwrap();
let result = serde_json::to_value(result).unwrap();
assert_eq!(result["tools"].as_array().unwrap().len(), 2);
assert_eq!(
result["_meta"]["litellm.ai/server_outcomes"]["alpha"]["tool_count"],
2
);
assert_eq!(
result["_meta"]["litellm.ai/server_outcomes"]["alpha-beta"]["status"],
"unreachable"
);
let prompts = harness
.gateway
.execute(Operation::ListPrompts(None), context(None))
.await
.unwrap();
assert_eq!(
serde_json::to_value(prompts).unwrap()["prompts"]
.as_array()
.unwrap()
.len(),
1
);
assert!(matches!(
harness
.gateway
.rest_tools(context(Some("alpha-beta")))
.await,
Err(Error::Upstream(_))
));
first.cancel().await.unwrap();
}
struct ResolvedToolPolicy;
impl litellm_gateway_mcp::OperationAuthorizer for ResolvedToolPolicy {
fn authorize<'a>(
&'a self,
server: &'a litellm_gateway_mcp::ServerInfo,
request: Option<&'a ClientRequest>,
) -> litellm_gateway_mcp::GatewayFuture<'a, ()> {
Box::pin(async move {
match request {
Some(ClientRequest::ListToolsRequest(_)) => Ok(()),
Some(ClientRequest::CallToolRequest(request))
if server.server_id == "id-alpha-beta"
&& request.params.name == "echo-value" =>
{
Ok(())
}
_ => Err(Error::Forbidden),
}
})
}
}
#[rstest]
#[case::resolved_alias("alpha-beta-echo-value", true)]
#[case::different_server("alpha-echo-value", false)]
#[case::different_tool("alpha-beta-fail", false)]
#[tokio::test]
async fn authorizes_resolved_server_and_tool_before_upstream_execution(
#[future] harness: Harness,
#[case] name: &str,
#[case] allowed: bool,
) {
let harness = harness.await;
let mut context = context(None);
context
.parts
.extensions
.insert(litellm_gateway_mcp::Authorization(std::sync::Arc::new(
ResolvedToolPolicy,
)));
let result = harness
.gateway
.execute(
Operation::CallTool(CallToolRequestParams::new(name.to_owned())),
context,
)
.await;
if allowed {
assert!(result.is_ok(), "{result:?}");
} else {
assert!(matches!(result, Err(Error::Forbidden)));
}
assert_eq!(harness.calls.load(Ordering::SeqCst), usize::from(allowed));
}
#[rstest]
#[tokio::test]
async fn initialize_checks_the_injected_operation_policy(#[future] harness: Harness) {
let harness = harness.await;
let mut context = context(Some("alpha"));
context
.parts
.extensions
.insert(litellm_gateway_mcp::Authorization(std::sync::Arc::new(
ResolvedToolPolicy,
)));
assert!(matches!(
harness.gateway.authorize(context).await,
Err(Error::Forbidden)
));
assert_eq!(harness.calls.load(Ordering::SeqCst), 0);
}
#[rstest]
#[case::mcp_allowed(true, true)]
#[case::mcp_denied(true, false)]
#[case::rest_allowed(false, true)]
#[case::rest_denied(false, false)]
#[tokio::test]
async fn http_transports_preserve_the_per_request_operation_policy(
#[future] harness: Harness,
#[case] protocol: bool,
#[case] allowed: bool,
) {
use axum::{
Extension,
body::{Body, to_bytes},
http::Request,
};
use tower::ServiceExt;
let harness = harness.await;
let server = if allowed { "alpha-beta" } else { "alpha" };
let (path, body) = if protocol {
(
"/mcp",
json!({"jsonrpc":"2.0", "id":1, "method":"tools/call", "params":{"name":format!("{server}-echo-value")}}),
)
} else {
(
"/mcp-rest/tools/call",
json!({"server_id":format!("id-{server}"), "name":"echo-value"}),
)
};
let app = harness
.app
.layer(Extension(litellm_gateway_mcp::Authorization(
std::sync::Arc::new(ResolvedToolPolicy),
)));
let response = app
.oneshot(
Request::post(path)
.header("host", "localhost")
.header("content-type", "application/json")
.header("accept", "application/json, text/event-stream")
.header(
"mcp-protocol-version",
ProtocolVersion::V_2025_11_25.as_str(),
)
.body(Body::from(body.to_string()))
.unwrap(),
)
.await
.unwrap();
assert_eq!(
response.status().as_u16(),
if protocol || allowed { 200 } else { 403 }
);
let wire = to_bytes(response.into_body(), 65536).await.unwrap();
if protocol {
let text = std::str::from_utf8(&wire).unwrap();
let json = text
.lines()
.find_map(|line| line.strip_prefix("data: "))
.unwrap();
let response: Value = serde_json::from_str(json).unwrap();
if allowed {
assert!(response.get("result").is_some(), "{response}");
} else {
assert_eq!(response["error"]["code"], -32003);
}
}
assert_eq!(harness.calls.load(Ordering::SeqCst), usize::from(allowed));
}

View file

@ -0,0 +1,131 @@
use litellm_gateway_mcp::RelayClient;
use rmcp::{
ClientHandler, ErrorData, RoleClient, RoleServer, ServerHandler, ServiceExt,
model::*,
service::{NotificationContext, RequestContext},
};
use rstest::rstest;
use serde_json::json;
use tokio::sync::mpsc;
struct Downstream(mpsc::UnboundedSender<ProgressNotificationParam>);
impl ClientHandler for Downstream {
fn get_info(&self) -> ClientConfig {
ClientConfig::new(
ClientCapabilities::builder().enable_elicitation().build(),
Implementation::new("fixture", "1"),
)
}
async fn create_elicitation(
&self,
_: ElicitRequestParams,
_: RequestContext<RoleClient>,
) -> Result<ElicitResult, ErrorData> {
Ok(
serde_json::from_value(json!({"action":"accept", "content":{"confirmed":true}}))
.unwrap(),
)
}
async fn on_progress(
&self,
params: ProgressNotificationParam,
_: NotificationContext<RoleClient>,
) {
self.0.send(params).unwrap();
}
}
struct AsksForInput;
impl ServerHandler for AsksForInput {
fn get_info(&self) -> ServerConfig {
ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
}
async fn call_tool(
&self,
_: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
let progress =
ProgressNotificationParam::new(context.meta.get_progress_token().unwrap(), 1.0);
context.peer.notify_progress(progress).await.unwrap();
let question = serde_json::from_value(json!({"message":"Confirm", "requestedSchema":{"type":"object","properties":{"confirmed":{"type":"boolean"}}}})).unwrap();
let answer = context.peer.create_elicitation(question).await.unwrap();
Ok(CallToolResult::success(vec![ContentBlock::text(
serde_json::to_string(&answer).unwrap(),
)])
.into())
}
}
struct Gateway;
impl ServerHandler for Gateway {
fn get_info(&self) -> ServerConfig {
ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
}
async fn call_tool(
&self,
request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
let (server, client) = tokio::io::duplex(65536);
let worker = tokio::spawn(async move {
AsksForInput
.serve(server)
.await
.unwrap()
.waiting()
.await
.unwrap()
});
let upstream = RelayClient::for_request(&context)
.serve(client)
.await
.unwrap();
let result = upstream.call_tool(request).await.unwrap();
upstream.cancel().await.unwrap();
worker.await.unwrap();
Ok(result.into())
}
}
#[rstest]
#[tokio::test]
async fn relays_elicitation_and_maps_progress_to_downstream_token() {
let (server, client) = tokio::io::duplex(65536);
let worker = tokio::spawn(async move {
Gateway
.serve(server)
.await
.unwrap()
.waiting()
.await
.unwrap()
});
let (sender, mut progress) = mpsc::unbounded_channel();
let client = Downstream(sender).serve(client).await.unwrap();
let request = CallToolRequest::new(CallToolRequestParams::new("confirm"));
let handle = client
.send_request_with_option(request.into(), Default::default())
.await
.unwrap();
let expected_token = handle.progress_token.clone();
let result = handle.await_response().await.unwrap();
let ServerResult::CallToolResult(result) = result else {
panic!("expected tool result")
};
let content: serde_json::Value =
serde_json::from_str(&result.content[0].as_text().unwrap().text).unwrap();
assert_eq!(content["content"]["confirmed"], true);
let update = progress.recv().await.unwrap();
assert_eq!(update.progress_token, expected_token);
assert_eq!(update.progress, 1.0);
client.cancel().await.unwrap();
worker.await.unwrap();
}

View file

@ -0,0 +1,194 @@
use std::{
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use axum::Router;
use litellm_gateway_mcp::{HttpConfig, NativeGateway, Registry, Server, ServerInfo, router};
use rmcp::{
ErrorData, RoleClient, RoleServer, ServerHandler, ServiceExt,
model::*,
service::{RequestContext, RunningService},
};
use rstest::fixture;
use serde_json::json;
pub struct Upstream {
pub calls: Arc<AtomicUsize>,
}
impl ServerHandler for Upstream {
fn get_info(&self) -> ServerConfig {
ServerConfig::new(
ServerCapabilities::builder()
.enable_tools()
.enable_prompts()
.enable_resources()
.build(),
)
}
async fn list_tools(
&self,
params: Option<PaginatedRequestParams>,
_: RequestContext<RoleServer>,
) -> Result<ListToolsResult, ErrorData> {
let second = params.and_then(|params| params.cursor).is_some();
let names = if second {
vec!["hidden", "fail"]
} else {
vec!["echo-value"]
};
let tools = names
.into_iter()
.map(|name| {
Tool::new(
name,
name,
Arc::new(serde_json::from_value(json!({"type":"object"})).unwrap()),
)
})
.collect();
let mut result = ListToolsResult::with_all_items(tools);
result.next_cursor = (!second).then(|| "page-2".into());
Ok(result)
}
async fn call_tool(
&self,
params: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
self.calls.fetch_add(1, Ordering::SeqCst);
if params.name == "fail" {
return Err(ErrorData::invalid_params("fixture tool failed", None));
}
Ok(CallToolResult::success(vec![ContentBlock::text(
json!({"name":params.name, "arguments":params.arguments, "meta":context.meta})
.to_string(),
)])
.into())
}
async fn list_prompts(
&self,
_: Option<PaginatedRequestParams>,
_: RequestContext<RoleServer>,
) -> Result<ListPromptsResult, ErrorData> {
Ok(ListPromptsResult::with_all_items(vec![Prompt::new(
"review-code",
None::<String>,
None,
)]))
}
async fn get_prompt(
&self,
params: GetPromptRequestParams,
_: RequestContext<RoleServer>,
) -> Result<GetPromptResponse, ErrorData> {
Ok(serde_json::from_value::<GetPromptResult>(
json!({"messages": [{"role":"user", "content":{"type":"text", "text":params.name}}]}),
)
.unwrap()
.into())
}
async fn list_resources(
&self,
_: Option<PaginatedRequestParams>,
_: RequestContext<RoleServer>,
) -> Result<ListResourcesResult, ErrorData> {
Ok(ListResourcesResult::with_all_items(vec![Resource::new(
"fixture://document",
"document",
)]))
}
async fn list_resource_templates(
&self,
_: Option<PaginatedRequestParams>,
_: RequestContext<RoleServer>,
) -> Result<ListResourceTemplatesResult, ErrorData> {
Ok(ListResourceTemplatesResult::with_all_items(vec![
serde_json::from_value(json!({"name":"files", "uriTemplate":"fixture://{file}"}))
.unwrap(),
]))
}
async fn read_resource(
&self,
params: ReadResourceRequestParams,
_: RequestContext<RoleServer>,
) -> Result<ReadResourceResponse, ErrorData> {
Ok(serde_json::from_value::<ReadResourceResult>(
json!({"contents":[{"uri":params.uri,"text":"fixture contents"}]}),
)
.unwrap()
.into())
}
}
pub struct Harness {
pub gateway: Arc<NativeGateway>,
pub app: Router,
pub calls: Arc<AtomicUsize>,
pub connections: Vec<RunningService<RoleClient, ()>>,
}
#[fixture]
pub async fn harness() -> Harness {
let calls = Arc::new(AtomicUsize::new(0));
let connections = futures_util::future::join_all((0..2).map(|_| {
let calls = calls.clone();
async move {
let (server, client) = tokio::io::duplex(65536);
tokio::spawn(async move {
Upstream { calls }
.serve(server)
.await
.unwrap()
.waiting()
.await
.unwrap();
});
().serve(client).await.unwrap()
}
}))
.await;
let servers: Vec<_> = connections
.iter()
.zip(["alpha", "alpha-beta"])
.map(|(connection, name)| Server {
info: ServerInfo {
server_id: format!("id-{name}"),
server_name: name.into(),
alias: None,
},
peer: connection.peer().clone(),
allowed_tools: Some(Arc::from(["echo-value".into(), "fail".into()])),
})
.collect();
let gateway = Arc::new(NativeGateway::new(
Arc::new(Registry::new(servers).unwrap()),
Duration::from_secs(2),
));
let app = router(gateway.clone(), HttpConfig::default());
Harness {
gateway,
app,
calls,
connections,
}
}
pub fn context(server: Option<&str>) -> litellm_gateway_mcp::Context {
litellm_gateway_mcp::Context {
parts: http::Request::new(()).into_parts().0,
server: server.map(str::to_owned),
mcp: None,
}
}

View file

@ -11,25 +11,28 @@ envy = "0.4.2"
http-body-util = "0.1"
litellm-core.workspace = true
litellm-gateway-inference.workspace = true
litellm-gateway-mcp.workspace = true
tokio-util = "0.7"
thiserror.workspace = true
futures-util.workspace = true
litellm-gateway-auth.workspace = true
litellm-gateway-ui.workspace = true
litellm-auth-types.workspace = true
litellm-gateway-auth.workspace = true
litellm-config.workspace = true
litellm-http.workspace = true
litellm-http = { workspace = true, features = ["mcp"] }
litellm-llms.workspace = true
litellm-secrets.workspace = true
litellm-tracing.workspace = true
serde_json.workspace = true
serde.workspace = true
tracing.workspace = true
tokio.workspace = true
tokio = { workspace = true, features = ["signal"] }
uuid.workspace = true
tower-sessions-moka-store = "0.15.0"
[dev-dependencies]
futures-util.workspace = true
rstest.workspace = true
base64.workspace = true
rstest.workspace = true
tempfile.workspace = true
tokio = { workspace = true, features = ["sync"] }
tower = { version = "0.5", features = ["util"] }

View file

@ -0,0 +1,43 @@
# Rust gateway
The binary reads YAML from `LITELLM_CONFIG` and mounts configured MCP servers alongside inference routes. MCP startup requires a resolvable `general_settings.master_key`. Every MCP and REST request uses the gateway's shared authentication middleware
## Run MCP locally
From the repository root, start the Python fixture in one terminal:
```sh
MCP_HOST=127.0.0.1 MCP_PORT=8090 .venv/bin/python tests/mcp_tests/mcp_e2e_upstream_server.py
```
Start the gateway in another terminal with the example configuration:
```sh
LITELLM_MASTER_KEY=local-mcp-key \
LITELLM_CONFIG=litellm-rust/crates/gateway/mcp_config.yaml \
HOST=127.0.0.1 PORT=4000 \
cargo run --manifest-path litellm-rust/Cargo.toml -p litellm-gateway
```
Call a tool through REST:
```sh
curl -sS http://127.0.0.1:4000/mcp-rest/tools/call \
-H 'Authorization: Bearer local-mcp-key' \
-H 'Content-Type: application/json' \
-d '{"server_id":"math","name":"add","arguments":{"a":2,"b":3}}'
```
MCP clients can connect to `/mcp` or the scoped `/math/mcp` endpoint with the same admission credential. Tool names on the MCP endpoint carry the configured alias or server-name prefix
## Configuration behavior
`Config::load` resolves nested `include` files once in breadth-first order. Included lists append and other top-level values replace earlier values. Paths resolve relative to the declaring file, with the root config directory as a fallback
`environment_variables` provides scalar overrides to the injected secret source and HTTP settings without changing process globals. Credentials, static headers, URLs, and stdio environment values can reference `os.environ/NAME`. Model providers and MCP share the host's HTTP settings, including TLS certificates and environment proxies. Explicit environment HTTP settings take precedence over `litellm_settings`
Each MCP entry supports HTTP or stdio transport, an optional pinned `server_id`, alias, tool allowlist, timeout in seconds, and positive `max_concurrent_requests`. Omitted or empty tool allowlists allow all upstream tools. HTTP entries accept `static_headers` and static API-key, bearer, basic, authorization, or token credentials. Stdio entries accept `command`, `args`, and `env`. Use explicit secret references in `env` to pass configured credentials to a child process
The gateway initializes upstreams before opening its listener. Startup fails when credentials cannot be resolved or an upstream cannot initialize. Shutdown cancels upstream connections and terminates owned stdio children
The config schema preserves Python sections, but parsing a section does not implement its service. MCP OAuth, database-backed server management, guardrails, spend accounting, local Python tools, and legacy SSE upstream connections remain unimplemented. Unsupported MCP settings and configured guardrail policies fail at startup. Incoming legacy SSE clients remain supported

View file

@ -0,0 +1,16 @@
model_list: []
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
mcp_allowed_hosts:
- localhost
- 127.0.0.1
mcp_servers:
math:
server_id: math
transport: http
url: http://127.0.0.1:8090/mcp
allowed_tools:
- add
- multiply
timeout: 30
max_concurrent_requests: 8

View file

@ -0,0 +1,73 @@
use std::sync::Arc;
use axum::{extract::Request, middleware::Next, response::Response};
use litellm_gateway_auth::{AccessRequest, AuthenticatedRequest, Error, McpAction};
use litellm_gateway_mcp::{
Authorization, GatewayFuture, OperationAuthorizer, ServerInfo, SessionOwner,
rmcp::model::ClientRequest,
};
pub(crate) async fn bind_session_owner(
identity: AuthenticatedRequest,
mut request: Request,
next: Next,
) -> Response {
request
.extensions_mut()
.insert(SessionOwner(identity.caller().session_owner()));
request
.extensions_mut()
.insert(Authorization(Arc::new(McpPolicy(identity))));
next.run(request).await
}
struct McpPolicy(AuthenticatedRequest);
impl OperationAuthorizer for McpPolicy {
fn authorize<'a>(
&'a self,
server: &'a ServerInfo,
request: Option<&'a ClientRequest>,
) -> GatewayFuture<'a, ()> {
Box::pin(async move {
let action = match request {
None => McpAction::Connect,
Some(ClientRequest::ListToolsRequest(_)) => McpAction::ListTools,
Some(ClientRequest::CallToolRequest(request)) => {
McpAction::CallTool(request.params.name.to_string())
}
Some(ClientRequest::ListPromptsRequest(_)) => McpAction::ListPrompts,
Some(ClientRequest::GetPromptRequest(request)) => {
McpAction::GetPrompt(request.params.name.clone())
}
Some(ClientRequest::ListResourcesRequest(_)) => McpAction::ListResources,
Some(ClientRequest::ListResourceTemplatesRequest(_)) => {
McpAction::ListResourceTemplates
}
Some(ClientRequest::ReadResourceRequest(request)) => {
McpAction::ReadResource(request.params.uri.clone())
}
Some(_) => return Err(litellm_gateway_mcp::Error::Forbidden),
};
let access = AccessRequest::Mcp {
server: server.server_id.clone(),
action,
};
self.0
.authorize(access.clone())
.await
.map_err(mcp_error)?
.consume(self.0.caller(), &access)
.map_err(mcp_error)
})
}
}
fn mcp_error(error: Error) -> litellm_gateway_mcp::Error {
match error {
Error::InvalidToken | Error::Expired | Error::Forbidden => {
litellm_gateway_mcp::Error::Forbidden
}
_ => litellm_gateway_mcp::Error::Configuration("MCP authorization unavailable".into()),
}
}

View file

@ -0,0 +1,13 @@
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error(transparent)]
Http(#[from] litellm_http::Error),
#[error("environment_variables values must be strings, numbers or booleans")]
Environment,
#[error("MCP host does not yet support configured setting {0}")]
McpSetting(String),
#[error(transparent)]
Mcp(#[from] litellm_gateway_mcp::ConnectError),
#[error("MCP requires a configured master key")]
Auth(#[from] litellm_gateway_auth::Error),
}

View file

@ -1,3 +1,8 @@
mod auth;
mod error;
mod secrets;
pub use error::Error;
use std::{sync::Arc, time::Instant};
use axum::{
@ -12,31 +17,112 @@ use http_body_util::BodyExt;
use litellm_config::Config;
use litellm_core::resources::CoreResources;
use litellm_gateway_auth::Auth;
use litellm_gateway_inference::{Gateway, ModelList};
use litellm_gateway_inference::{Gateway, ModelRouter};
use litellm_http::{
ClientVariant, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
ClientVariant, HttpClientPool, HttpSettings, HttpSettingsLayer, Resolution, SslVerify,
media::PublicDnsResolver,
};
use litellm_secrets::source::EnvironmentSecrets;
use litellm_tracing::ByteChunk;
use uuid::Uuid;
pub fn build_inference(config: &Config) -> Result<Arc<Gateway>, litellm_http::Error> {
pub fn build_inference(config: &Config) -> Result<Arc<Gateway>, Error> {
let pool = Arc::new(HttpClientPool::new(Arc::new(PublicDnsResolver)));
let http = Resolution::from(&HttpSettings::default()).config;
let environment = secrets::environment_values(&config.environment_variables)?;
let lookup = |name: &str| {
environment
.get(name)
.map(|value| value.expose().to_owned())
.or_else(|| std::env::var(name).ok())
};
let settings = &config.litellm_settings;
let http_settings = HttpSettings::from_layers([
HttpSettingsLayer::from_environment(&lookup),
HttpSettingsLayer {
ssl_verify: settings.ssl_verify.as_ref().map(|value| match value {
litellm_config::Flag::Boolean(true) => SslVerify::Enabled,
litellm_config::Flag::Boolean(false) => SslVerify::Disabled,
litellm_config::Flag::String(value) => SslVerify::parse(value),
}),
ssl_certificate: settings.ssl_certificate.as_ref().map(Into::into),
ssl_security_level: settings.ssl_security_level.clone(),
ssl_ecdh_curve: settings.ssl_ecdh_curve.clone(),
force_ipv4: settings.force_ipv4,
http2: settings.http2,
aiohttp_trust_env: settings.aiohttp_trust_env,
disable_aiohttp_trust_env: settings.disable_aiohttp_trust_env,
disable_aiohttp_transport: settings.disable_aiohttp_transport,
..Default::default()
},
]);
let http = Resolution::from(&http_settings).config;
let client = pool.client(&http, ClientVariant::Provider)?;
let secrets = Arc::new(EnvironmentSecrets::python_compatible(client));
let secrets = Arc::new(secrets::ConfigSecrets::new(
environment,
Arc::new(EnvironmentSecrets::python_compatible(client)),
));
let resources = CoreResources::new(pool);
Ok(Arc::new(Gateway::new(
resources,
http,
secrets,
ModelList::from_model_list(&config.model_list),
ModelRouter::from_model_list(&config.model_list),
)?))
}
pub fn router(inference: Arc<Gateway>, config: &Config, ui: Option<Router>) -> Router {
pub async fn build_mcp(
config: &Config,
secrets: Arc<dyn litellm_secrets::source::SecretSource>,
shutdown: tokio_util::sync::CancellationToken,
pool: &HttpClientPool,
http: &litellm_http::HttpClientConfig,
) -> Result<Option<litellm_gateway_mcp::ConfiguredGateway>, Error> {
if !config.mcp_tools.is_empty() {
return Err(Error::McpSetting("mcp_tools".into()));
}
if config.mcp_servers.is_empty() {
return Ok(None);
}
if let Some(setting) = config
.general_settings
.additional_fields
.keys()
.chain(config.litellm_settings.additional_fields.keys())
.find(|key| key.starts_with("mcp_"))
{
return Err(Error::McpSetting(setting.clone()));
}
if !config.guardrails.is_empty()
|| !config.policies.is_empty()
|| !config.policy_attachments.is_empty()
{
return Err(Error::McpSetting("guardrails or policies".into()));
}
Auth::from_config(config, secrets.clone())
.validate()
.await?;
let client = pool.mcp_client(http)?;
Ok(Some(
litellm_gateway_mcp::ConfiguredGateway::connect(
&config.mcp_servers,
client,
secrets.as_ref(),
shutdown,
)
.await?,
))
}
pub fn router(
inference: Arc<Gateway>,
config: &Config,
ui: Option<Router>,
mcp: Option<Router>,
) -> Router {
let auth = Auth::from_config(config, inference.secrets.clone());
let inference = litellm_gateway_inference::router(inference)
.merge(mcp.unwrap_or_default())
.route_layer(axum::middleware::from_fn(auth::bind_session_owner))
.route_layer(axum::middleware::from_fn_with_state(
auth,
litellm_gateway_auth::authenticate,

View file

@ -63,10 +63,70 @@ async fn main() -> Result<(), Box<dyn Error>> {
}
None => None,
};
let shutdown = tokio_util::sync::CancellationToken::new();
let _shutdown_guard = shutdown.clone().drop_guard();
let signal_shutdown = shutdown.clone();
tokio::spawn(async move {
shutdown_signal().await;
signal_shutdown.cancel();
});
let mcp = litellm_gateway::build_mcp(
&config,
inference.secrets.clone(),
shutdown.clone(),
&inference.resources.pool,
&inference.http,
)
.await?;
let mcp_router = mcp.as_ref().map(|gateway| {
let defaults = litellm_gateway_mcp::HttpConfig::default();
litellm_gateway_mcp::router(
gateway.operations.clone(),
litellm_gateway_mcp::HttpConfig {
allowed_hosts: config
.general_settings
.mcp_allowed_hosts
.as_deref()
.map(Vec::from)
.unwrap_or(defaults.allowed_hosts),
allowed_origins: config.general_settings.mcp_allowed_origins.to_vec(),
cancellation_token: shutdown.clone(),
server_info: defaults.server_info,
},
)
});
let listener = tokio::net::TcpListener::bind((settings.host.as_str(), settings.port)).await?;
tracing::info!(address = %listener.local_addr()?, models = config.model_list.len(), log_level = %level, "gateway listening");
axum::serve(listener, litellm_gateway::router(inference, &config, ui)).await?;
let result = axum::serve(
listener,
litellm_gateway::router(inference, &config, ui, mcp_router),
)
.with_graceful_shutdown(shutdown.clone().cancelled_owned())
.await;
shutdown.cancel();
if let Some(mcp) = mcp {
mcp.close().await;
}
result?;
Ok(())
}
async fn shutdown_signal() {
#[cfg(unix)]
{
match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
Ok(mut terminate) => {
tokio::select! { _ = tokio::signal::ctrl_c() => (), _ = terminate.recv() => () }
}
Err(_) => {
let _ = tokio::signal::ctrl_c().await;
}
}
}
#[cfg(not(unix))]
{
let _ = tokio::signal::ctrl_c().await;
}
}

View file

@ -0,0 +1,48 @@
use std::{collections::BTreeMap, sync::Arc};
use futures_util::future::BoxFuture;
use litellm_auth_types::SecretValue;
use litellm_config::{Object, Value};
use litellm_secrets::source::SecretSource;
use crate::Error;
pub(super) struct ConfigSecrets {
values: BTreeMap<String, SecretValue>,
fallback: Arc<dyn SecretSource>,
}
impl ConfigSecrets {
pub fn new(values: BTreeMap<String, SecretValue>, fallback: Arc<dyn SecretSource>) -> Self {
Self { values, fallback }
}
}
impl SecretSource for ConfigSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
Box::pin(async move {
match self.values.get(name) {
Some(value) => Ok(Some(value.clone())),
None => self.fallback.get_secret_str(name).await,
}
})
}
}
pub(super) fn environment_values(values: &Object) -> Result<BTreeMap<String, SecretValue>, Error> {
values
.iter()
.map(|(key, value)| {
let text = match value {
Value::String(value) => value.clone(),
Value::Number(value) => value.to_string(),
Value::Bool(value) => value.to_string(),
_ => return Err(Error::Environment),
};
Ok((key.clone(), SecretValue::new(text)))
})
.collect()
}

View file

@ -0,0 +1,344 @@
use std::{sync::Arc, time::Duration};
use axum::{
Router,
body::{Body, to_bytes},
http::Request,
};
use litellm_config::Config;
use litellm_gateway_mcp::{
HttpConfig,
rmcp::{
ErrorData, RoleServer, ServerHandler,
model::*,
service::RequestContext,
transport::streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
},
},
};
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use tokio_util::sync::CancellationToken;
use tower::ServiceExt;
#[derive(Clone)]
struct Upstream;
impl ServerHandler for Upstream {
fn get_info(&self) -> ServerConfig {
ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
}
async fn list_tools(
&self,
_: Option<PaginatedRequestParams>,
_: RequestContext<RoleServer>,
) -> Result<ListToolsResult, ErrorData> {
Ok(ListToolsResult::with_all_items(
["echo", "hidden"]
.into_iter()
.map(|name| {
Tool::new(
name,
name,
Arc::new(serde_json::from_value(json!({"type":"object"})).unwrap()),
)
})
.collect(),
))
}
async fn call_tool(
&self,
params: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> Result<CallToolResponse, ErrorData> {
let headers = &context
.extensions
.get::<axum::http::request::Parts>()
.unwrap()
.headers;
let result = json!({"arguments": params.arguments, "authorization": headers.get("authorization").unwrap().to_str().unwrap(), "static": headers.get("x-static").unwrap().to_str().unwrap()});
Ok(CallToolResult::success(vec![ContentBlock::text(result.to_string())]).into())
}
}
struct Fixture {
url: String,
shutdown: CancellationToken,
task: tokio::task::JoinHandle<()>,
}
#[fixture]
async fn upstream() -> Fixture {
let shutdown = CancellationToken::new();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/mcp", listener.local_addr().unwrap());
let transport = StreamableHttpService::new(
|| Ok(Upstream),
Arc::new(LocalSessionManager::default()),
StreamableHttpServerConfig::default().with_cancellation_token(shutdown.clone()),
);
let app = Router::new().route_service("/mcp", transport);
let stop = shutdown.clone();
let task = tokio::spawn(async move {
axum::serve(listener, app)
.with_graceful_shutdown(stop.cancelled_owned())
.await
.unwrap();
});
Fixture {
url,
shutdown,
task,
}
}
#[rstest]
#[tokio::test]
async fn serves_configured_upstream_with_separate_admission_and_outbound_credentials(
#[future] upstream: Fixture,
) {
let upstream = upstream.await;
let config = Config::from_yaml(&format!(
r#"
environment_variables:
GATEWAY_KEY: admission-secret
UPSTREAM_KEY: upstream-secret
STATIC_VALUE: static-secret
general_settings:
master_key: os.environ/GATEWAY_KEY
mcp_servers:
docs:
server_id: configured-id
alias: knowledge
transport: http
url: {}
auth_type: bearer_token
authentication_token: os.environ/UPSTREAM_KEY
static_headers:
x-static: os.environ/STATIC_VALUE
allowed_tools: [echo]
timeout: 2
max_concurrent_requests: 1
"#,
upstream.url
))
.unwrap();
let inference = litellm_gateway::build_inference(&config).unwrap();
let shutdown = CancellationToken::new();
let mcp = litellm_gateway::build_mcp(
&config,
inference.secrets.clone(),
shutdown.clone(),
&inference.resources.pool,
&inference.http,
)
.await
.unwrap()
.unwrap();
let mcp_router = litellm_gateway_mcp::router(
mcp.operations.clone(),
HttpConfig {
cancellation_token: shutdown.clone(),
..Default::default()
},
);
let app = litellm_gateway::router(inference, &config, None, Some(mcp_router));
for (path, method) in [
("/mcp", "POST"),
("/mcp", "GET"),
("/mcp", "DELETE"),
("/mcp/sse", "GET"),
("/mcp/sse/messages", "POST"),
("/mcp-rest/tools/list", "GET"),
("/mcp/enabled", "GET"),
] {
let response = app
.clone()
.oneshot(
Request::builder()
.method(method)
.uri(path)
.header("host", "localhost")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), 401, "{method} {path}");
}
let list = app
.clone()
.oneshot(
Request::builder()
.uri("/mcp-rest/tools/list?server_id=configured-id")
.header("authorization", "Bearer admission-secret")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(list.status(), 200);
let body: Value =
serde_json::from_slice(&to_bytes(list.into_body(), 65536).await.unwrap()).unwrap();
assert_eq!(body["tools"].as_array().unwrap().len(), 1);
assert_eq!(body["tools"][0]["name"], "echo");
assert_eq!(body["tools"][0]["mcp_info"]["alias"], "knowledge");
let response = app.clone().oneshot(Request::post("/mcp-rest/tools/call").header("authorization", "Bearer admission-secret").header("content-type", "application/json").body(Body::from(json!({"server_id":"configured-id", "name":"echo", "arguments":{"value":"hello"}}).to_string())).unwrap()).await.unwrap();
assert_eq!(response.status(), 200);
let body: Value =
serde_json::from_slice(&to_bytes(response.into_body(), 65536).await.unwrap()).unwrap();
let result: Value = serde_json::from_str(body["content"][0]["text"].as_str().unwrap()).unwrap();
assert_eq!(
result,
json!({"arguments":{"value":"hello"}, "authorization":"Bearer upstream-secret", "static":"static-secret"})
);
let blocked = app
.oneshot(
Request::post("/mcp-rest/tools/call")
.header("authorization", "Bearer admission-secret")
.header("content-type", "application/json")
.body(Body::from(
json!({"server_id":"configured-id", "name":"hidden"}).to_string(),
))
.unwrap(),
)
.await
.unwrap();
assert_eq!(blocked.status(), 403);
shutdown.cancel();
tokio::time::timeout(Duration::from_secs(5), mcp.close())
.await
.unwrap();
upstream.shutdown.cancel();
tokio::time::timeout(Duration::from_secs(5), upstream.task)
.await
.unwrap()
.unwrap();
}
#[rstest]
#[case::unsupported_auth("auth_type: oauth2", "auth mode")]
#[case::unsupported_policy("allowed_params: {echo: [safe]}", "allowed_params")]
#[case::unsupported_transport("transport: sse", "legacy SSE")]
#[case::zero_timeout("timeout: 0", "timeout")]
#[case::negative_timeout("timeout: -1", "timeout")]
#[case::zero_concurrency("max_concurrent_requests: 0", "max_concurrent_requests")]
#[case::protocol_override("mcp_info: {protocol_version: auto}", "protocol and cost")]
#[case::costs("mcp_info: {mcp_server_cost_info: {default: 1}}", "protocol and cost")]
#[case::missing_token("auth_type: api_key", "authentication_token")]
#[case::missing_env(
"auth_type: api_key\n authentication_token: os.environ/LITELLM_TEST_NONEXISTENT_MCP_SECRET",
"referenced secret"
)]
#[case::reserved_header("static_headers: {host: secret-value}", "header")]
#[tokio::test]
async fn rejects_unsupported_or_invalid_settings_before_serving(
#[case] settings: &str,
#[case] expected: &str,
) {
let config = Config::from_yaml(&format!("general_settings: {{master_key: gateway-key}}\nmcp_servers:\n docs:\n url: http://127.0.0.1:1/mcp\n {settings}\n")).unwrap();
let inference = litellm_gateway::build_inference(&config).unwrap();
let result = litellm_gateway::build_mcp(
&config,
Arc::new(NoSecrets),
CancellationToken::new(),
&inference.resources.pool,
&inference.http,
)
.await;
let error = result.err().expect("startup must fail").to_string();
assert!(error.contains(expected), "{error}");
assert!(!error.contains("secret-value"));
}
#[rstest]
#[tokio::test]
async fn mcp_is_optional_and_requires_admission_credentials_when_enabled() {
let empty = Config::from_yaml("{}").unwrap();
let inference = litellm_gateway::build_inference(&empty).unwrap();
assert!(
litellm_gateway::build_mcp(
&empty,
inference.secrets.clone(),
CancellationToken::new(),
&inference.resources.pool,
&inference.http
)
.await
.unwrap()
.is_none()
);
let enabled =
Config::from_yaml("mcp_servers: {docs: {url: 'http://127.0.0.1:1/mcp'}}").unwrap();
let error = litellm_gateway::build_mcp(
&enabled,
inference.secrets.clone(),
CancellationToken::new(),
&inference.resources.pool,
&inference.http,
)
.await
.err()
.unwrap();
assert!(matches!(error, litellm_gateway::Error::Auth(_)));
}
struct NoSecrets;
impl litellm_secrets::source::SecretSource for NoSecrets {
fn get_secret_str<'a>(
&'a self,
_: &'a str,
) -> futures_util::future::BoxFuture<
'a,
Result<Option<litellm_auth_types::SecretValue>, litellm_secrets::Error>,
> {
Box::pin(async { Ok(None) })
}
}
#[rstest]
fn applies_configured_http_environment_before_building_clients() {
let directory = tempfile::tempdir().unwrap();
let missing = directory.path().join("missing-ca.pem");
let config = Config::from_yaml(&format!(
"environment_variables: {{SSL_VERIFY: '{}'}}\nlitellm_settings: {{ssl_verify: false}}",
missing.display(),
))
.unwrap();
let error = litellm_gateway::build_inference(&config).err().unwrap();
assert!(matches!(
error,
litellm_gateway::Error::Http(litellm_http::Error::Read { path, .. })
if path == missing
));
}
#[rstest]
#[case::client_policy("general_settings: {master_key: key, mcp_allowed_clients: [approved]}")]
#[case::guardrail("guardrails: [{guardrail_name: policy}]")]
#[case::local_tools("mcp_tools: [{name: custom}]")]
#[tokio::test]
async fn refuses_unimplemented_global_mcp_policy(#[case] settings: &str) {
let config = Config::from_yaml(&format!(
"{settings}\nmcp_servers: {{docs: {{url: 'http://127.0.0.1:1/mcp'}}}}\n"
))
.unwrap();
let inference = litellm_gateway::build_inference(&config).unwrap();
let error = litellm_gateway::build_mcp(
&config,
Arc::new(NoSecrets),
CancellationToken::new(),
&inference.resources.pool,
&inference.http,
)
.await
.err()
.unwrap();
assert!(matches!(error, litellm_gateway::Error::McpSetting(_)));
}

View file

@ -72,11 +72,14 @@ async fn authenticates_before_serving_mounted_inference_routes(
let address = listener.local_addr().unwrap();
let (shutdown, stopped) = oneshot::channel();
let server = tokio::spawn(async move {
axum::serve(listener, litellm_gateway::router(inference, &config, None))
.with_graceful_shutdown(async move {
let _ = stopped.await;
})
.await
axum::serve(
listener,
litellm_gateway::router(inference, &config, None, None),
)
.with_graceful_shutdown(async move {
let _ = stopped.await;
})
.await
});
let request = client
@ -126,7 +129,7 @@ async fn logs_request_outcome_without_credentials_or_query(inference: Arc<Gatewa
let logger = Logger::new(LogSink(sender));
let response = logger
.instrument(litellm_gateway::router(inference, &config, None).oneshot(request))
.instrument(litellm_gateway::router(inference, &config, None, None).oneshot(request))
.await
.unwrap();
@ -162,7 +165,7 @@ async fn mounts_ui_without_exposing_credentials_or_authorizing_inference(inferen
false,
)
.merge(litellm_gateway_ui::dashboard_assets(assets.path()));
let app = litellm_gateway::router(inference, &config, Some(ui));
let app = litellm_gateway::router(inference, &config, Some(ui), None);
let page = app
.clone()
.oneshot(Request::get("/ui/").body(Body::empty()).unwrap())
@ -272,7 +275,7 @@ async fn ui_routes_are_absent_when_not_mounted(
let config =
Config::from_yaml("model_list: []\ngeneral_settings:\n master_key: gateway-key\n")
.unwrap();
let response = litellm_gateway::router(inference, &config, None)
let response = litellm_gateway::router(inference, &config, None, None)
.oneshot(
Request::builder()
.method(method)

View file

@ -7,8 +7,11 @@ repository.workspace = true
[features]
test-support = []
mcp = ["dep:rmcp", "dep:reqwest-mcp"]
[dependencies]
rmcp = { version = "=3.4.1", default-features = false, features = ["client", "transport-streamable-http-client-reqwest"], optional = true }
reqwest-mcp = { package = "reqwest", version = "0.13.2", default-features = false, features = ["rustls-no-provider", "json", "stream", "http2"], optional = true }
http.workspace = true
litellm-core-utils.workspace = true
hyper-util.workspace = true

View file

@ -7,6 +7,8 @@
mod client;
mod config;
mod error;
#[cfg(feature = "mcp")]
mod mcp;
pub mod media;
pub mod outbound;
mod pool;

View file

@ -0,0 +1,76 @@
use std::{
collections::HashMap,
net::{IpAddr, Ipv4Addr},
sync::{Mutex, PoisonError},
time::{Duration, Instant},
};
use crate::{Error, HttpClientConfig};
struct Entry {
client: reqwest_mcp::Client,
built_at: Instant,
}
#[derive(Default)]
pub(super) struct Pool(Mutex<HashMap<HttpClientConfig, Entry>>);
impl Pool {
pub fn client(
&self,
config: &HttpClientConfig,
ttl: Duration,
) -> Result<reqwest_mcp::Client, Error> {
let mut clients = self.0.lock().unwrap_or_else(PoisonError::into_inner);
if let Some(entry) = clients.get(config)
&& entry.built_at.elapsed() < ttl
{
return Ok(entry.client.clone());
}
let client = build(config)?;
clients.insert(
config.clone(),
Entry {
client: client.clone(),
built_at: Instant::now(),
},
);
Ok(client)
}
}
fn build(config: &HttpClientConfig) -> Result<reqwest_mcp::Client, Error> {
let base = reqwest_mcp::Client::builder()
.tls_backend_preconfigured(rustls::ClientConfig::try_from(config)?)
.connect_timeout(config.connect_timeout)
.pool_idle_timeout(config.pool_idle_timeout)
.redirect(reqwest_mcp::redirect::Policy::none());
let keepalive = match config.tcp_keepalive {
Some(value) => base
.tcp_keepalive(value.idle)
.tcp_keepalive_interval(value.interval)
.tcp_keepalive_retries(value.retries),
None => base,
};
let address = if config.force_ipv4 {
keepalive.local_address(IpAddr::V4(Ipv4Addr::UNSPECIFIED))
} else {
keepalive
};
let protocol = if config.http2 {
address
} else {
address.http1_only()
};
let agent = match &config.user_agent {
Some(value) => protocol.user_agent(value),
None => protocol,
};
config
.proxies
.mcp_proxies()
.into_iter()
.fold(agent.no_proxy(), reqwest_mcp::ClientBuilder::proxy)
.build()
.map_err(|error| Error::Client(error.without_url().to_string()))
}

View file

@ -29,6 +29,8 @@ pub struct HttpClientPool {
media_resolver: Arc<dyn Resolve>,
ttl: Duration,
clients: Mutex<Clients>,
#[cfg(feature = "mcp")]
mcp: crate::mcp::Pool,
}
impl HttpClientPool {
@ -41,9 +43,19 @@ impl HttpClientPool {
media_resolver,
ttl,
clients: Mutex::default(),
#[cfg(feature = "mcp")]
mcp: crate::mcp::Pool::default(),
}
}
#[cfg(feature = "mcp")]
pub fn mcp_client(
&self,
config: &HttpClientConfig,
) -> Result<impl rmcp::transport::streamable_http_client::StreamableHttpClient, Error> {
self.mcp.client(config, self.ttl)
}
pub fn client(
&self,
config: &HttpClientConfig,
@ -378,4 +390,56 @@ mod tests {
.is_err()
);
}
#[cfg(feature = "mcp")]
#[rstest::rstest]
#[tokio::test]
async fn mcp_clients_reuse_connections_and_apply_proxy_headers_without_redirects() {
let (address, connections, requests) = serve("HTTP/1.1 302 Found").await;
let pool = pool();
let config = HttpClientConfig {
proxies: proxied_through(&format!("http://{address}")),
..config("mcp-pool-test")
};
for _ in 0..2 {
let response = pool
.mcp
.client(&config, pool.ttl)
.unwrap()
.get("http://upstream.invalid/mcp")
.timeout(Duration::from_secs(2))
.send()
.await
.unwrap();
assert_eq!(response.status(), 302);
response.bytes().await.unwrap();
}
assert_eq!(connections.load(Ordering::SeqCst), 1);
let requests = requests.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(
requests
.iter()
.all(|request| request.starts_with("GET http://upstream.invalid/mcp HTTP/1.1"))
);
assert!(
requests
.iter()
.all(|request| request.contains("user-agent: mcp-pool-test"))
);
}
#[cfg(feature = "mcp")]
#[rstest::rstest]
#[tokio::test]
async fn mcp_client_uses_host_tls_configuration() {
let directory = tempfile::tempdir().unwrap();
let config = HttpClientConfig {
verify: Verify::CaBundle(directory.path().join("missing.pem")),
..config("mcp-test")
};
assert!(matches!(
pool().mcp_client(&config),
Err(Error::Read { .. })
));
}
}

View file

@ -63,6 +63,19 @@ impl EnvironmentProxies {
.map(|proxy| proxy.no_proxy(no_proxy.clone()))
.collect()
}
#[cfg(feature = "mcp")]
pub(crate) fn mcp_proxies(&self) -> Vec<reqwest_mcp::Proxy> {
let no_proxy = reqwest_mcp::NoProxy::from_string(&self.no);
[
reqwest_mcp::Proxy::http(self.http.as_str()),
reqwest_mcp::Proxy::https(self.https.as_str()),
reqwest_mcp::Proxy::all(self.all.as_str()),
]
.into_iter()
.filter_map(Result::ok)
.map(|proxy| proxy.no_proxy(no_proxy.clone()))
.collect()
}
}
#[cfg(test)]

View file

@ -18,7 +18,7 @@ For Mistral, `async_transform_ocr_request` uses the base default in both languag
For non-OCR pairs, order corresponding methods as parameter support/mapping, environment validation, URL construction, request transformation, and response transformation, followed by Rust-only runtime hooks. Auth resolution remains split between configs and route preparation in litellm-core. Chat `supported_openai_param_mappings` describes accepted OpenAI/provider name pairs, unlike Python's `get_supported_openai_params` name list. Audio `map_transcription_params` remains a Rust filtering helper
Azure Messages maps to `llms/azure_ai/anthropic/messages_transformation.py`; Bedrock Converse maps to `llms/bedrock/chat/converse_transformation.py`. `AnthropicConfig`, `AmazonConverseConfig`, and the non-OCR base traits are partial ports. `OpenAiResponsesApiConfig` currently implements only the WebSocket surface. Preserve their acceptance gates, passthrough behavior, and host fallback contracts when aligning layout
Azure Messages maps to `llms/azure_ai/anthropic/messages_transformation.py`; Bedrock Converse maps to `llms/bedrock/chat/converse_transformation.py`. `AnthropicConfig`, `AmazonConverseConfig`, and the non-OCR base traits are partial ports. `OpenAiResponsesApiConfig` implements WebSocket transformations and a direct HTTP Responses path. Its HTTP path does not implement Python model-specific parameter rewriting or Responses-to-Chat emulation. Preserve their acceptance gates, passthrough behavior, and host fallback contracts when aligning layout
## Provider and format boundaries