diff --git a/.github/dependabot.yml b/.github/dependabot.yml index 0fa34ce69..168b680a6 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -21,6 +21,10 @@ updates: # Mark PRs as CI related change. - T-CI open-pull-requests-limit: 3 + groups: + codeql-action: + patterns: + - "github/codeql-action*" commit-message: prefix: "chore" include: "scope" diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a43adf760..1edfe69b2 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -31,12 +31,10 @@ jobs: node-version: '22' - name: Install commitlint - run: | - npm install --save-dev @commitlint/cli @commitlint/config-conventional - echo "module.exports = {extends: ['@commitlint/config-conventional']}" > commitlint.config.js - + run: npm install --global @commitlint/cli@20.4.3 @commitlint/config-conventional@20.4.3 + - name: Lint commit messages - run: npx commitlint --from ${{ github.event.pull_request.base.sha }} --to ${{ github.event.pull_request.head.sha }} --verbose + run: commitlint --extends @commitlint/config-conventional --from ${{ github.event.pull_request.base.sha }} --to ${{ github.event.pull_request.head.sha }} --verbose fmt: name: Code Formatting @@ -61,9 +59,17 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - - name: Run clippy + - name: Run clippy (all features) run: cargo clippy --all-targets --all-features -- -D warnings + - name: Run clippy (all features except local) + run: | + FEATURES=$(cargo metadata --no-deps --format-version 1 \ + | jq -r '[.packages[] | select(.name == "rmcp") | .features | keys[] + | select(startswith("__") | not) + | select(. != "local")] | join(",")') + cargo clippy --package rmcp --all-targets --no-default-features --features "$FEATURES" -- -D warnings + semver: name: SemVer Check runs-on: ubuntu-latest @@ -79,7 +85,7 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-semver-checks - uses: taiki-e/install-action@b6ff580856c41316412a0b9b60540fbc6f8c82cc # v2.86.7 + uses: taiki-e/install-action@7b8d4719ee4aaa279bdf55df38dacb9ebfe12a6c # v2.87.6 with: tool: cargo-semver-checks @@ -133,7 +139,7 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-public-api - uses: taiki-e/install-action@b6ff580856c41316412a0b9b60540fbc6f8c82cc # v2.86.7 + uses: taiki-e/install-action@7b8d4719ee4aaa279bdf55df38dacb9ebfe12a6c # v2.87.6 with: tool: cargo-public-api @@ -206,7 +212,7 @@ jobs: steps: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Spell Check Repo - uses: crate-ci/typos@1a51d4b5a03bb97576af186c813af67e9137ba7c # master + uses: crate-ci/typos@d43b6c087ac471e2ea7b8af622ff15f05c0c365b # master msrv: name: Check MSRV @@ -245,7 +251,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable @@ -274,7 +280,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable @@ -310,7 +316,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable @@ -345,7 +351,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index 4ecf3484b..83e1ce4d8 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -27,13 +27,13 @@ jobs: uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Initialize CodeQL - uses: github/codeql-action/init@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 + uses: github/codeql-action/init@cdf488f595d80d6e07e03d4674febd5ab45fa938 # v4.37.9 with: languages: ${{ matrix.language }} config-file: ./.github/codeql/codeql-config.yml - name: Autobuild - uses: github/codeql-action/autobuild@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 + uses: github/codeql-action/autobuild@cdf488f595d80d6e07e03d4674febd5ab45fa938 # v4.37.9 - name: Perform CodeQL Analysis - uses: github/codeql-action/analyze@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 + uses: github/codeql-action/analyze@cdf488f595d80d6e07e03d4674febd5ab45fa938 # v4.37.9 diff --git a/Cargo.toml b/Cargo.toml index 001096245..261b77e83 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,13 +4,13 @@ default-members = ["crates/rmcp", "crates/rmcp-macros"] resolver = "2" [workspace.dependencies] -rmcp = { version = "3.1.4", path = "./crates/rmcp" } -rmcp-macros = { version = "3.1.4", path = "./crates/rmcp-macros" } +rmcp = { version = "3.3.0", path = "./crates/rmcp" } +rmcp-macros = { version = "3.3.0", path = "./crates/rmcp-macros" } [workspace.package] edition = "2024" rust-version = "1.88" -version = "3.1.4" +version = "3.3.0" authors = ["4t145 "] license = "Apache-2.0" repository = "https://github.com/modelcontextprotocol/rust-sdk/" diff --git a/conformance/Cargo.toml b/conformance/Cargo.toml index fc4be6d90..e8de6f072 100644 --- a/conformance/Cargo.toml +++ b/conformance/Cargo.toml @@ -19,6 +19,7 @@ rmcp = { path = "../crates/rmcp", features = [ "elicitation", "auth", "auth-client-credentials-jwt", + "auth-enterprise-managed", "request-state", "transport-streamable-http-server", "transport-streamable-http-client-reqwest", @@ -31,6 +32,7 @@ tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter"] } axum = { version = "0.8", features = ["macros"] } anyhow = "1" +oauth2 = { version = "5.0", default-features = false } reqwest = { version = "0.13", features = ["json"] } urlencoding = "2" url = "2" diff --git a/conformance/src/bin/client.rs b/conformance/src/bin/client.rs index 5d654105a..92e4a945a 100644 --- a/conformance/src/bin/client.rs +++ b/conformance/src/bin/client.rs @@ -1,3 +1,10 @@ +#![expect( + deprecated, + reason = "The conformance suite still exercises deprecated sampling scenarios" +)] + +use anyhow::Context; +use oauth2::{ClientSecret, RefreshToken}; use rmcp::{ ClientHandler, ClientLifecycleMode, ClientServiceExt, ErrorData, RoleClient, ServiceExt, model::*, @@ -6,7 +13,8 @@ use rmcp::{ AuthClient, AuthorizationManager, StreamableHttpClientTransport, auth::{ AuthorizationCallback, AuthorizationRequest, ClientCredentialsConfig, - InMemoryCredentialStore, JwtSigningAlgorithm, OAuthState, + InMemoryCredentialStore, JwtSigningAlgorithm, OAuthState, default_oauth_http_client, + enterprise::{EmaAuthorizationServer, EmaClientAuthentication, EmaExchangeRequest}, }, streamable_http_client::StreamableHttpClientTransportConfig, }, @@ -36,6 +44,17 @@ struct ConformanceContext { private_key_pem: Option, #[serde(default)] signing_algorithm: Option, + // enterprise-managed-authorization-refresh-token + #[serde(default)] + idp_client_id: Option, + #[serde(default)] + idp_client_secret: Option, + #[serde(default)] + idp_refresh_token: Option, + #[serde(default)] + idp_issuer: Option, + #[serde(default)] + idp_token_endpoint: Option, } fn load_context() -> ConformanceContext { @@ -760,6 +779,66 @@ async fn run_client_credentials_jwt( Ok(()) } +/// Exchange the fixture's IdP refresh token, then exercise authenticated MCP access. +async fn run_ema_refresh_token_client( + server_url: &str, + ctx: &ConformanceContext, +) -> anyhow::Result<()> { + let manager = AuthorizationManager::new(server_url).await?; + let metadata = manager.resolve_metadata().await?.metadata; + let idp = EmaAuthorizationServer::new( + ctx.idp_issuer.as_deref().context("Missing idp_issuer")?, + ctx.idp_token_endpoint + .as_deref() + .context("Missing idp_token_endpoint")?, + ctx.idp_client_id + .as_deref() + .context("Missing idp_client_id")?, + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic( + ClientSecret::new( + ctx.idp_client_secret + .clone() + .context("Missing idp_client_secret")?, + ), + )); + let resource_as = EmaAuthorizationServer::new( + metadata + .issuer + .context("Missing authorization server issuer")?, + metadata.token_endpoint, + ctx.client_id.as_deref().context("Missing client_id")?, + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic( + ClientSecret::new(ctx.client_secret.clone().context("Missing client_secret")?), + )); + let refresh_token = RefreshToken::new( + ctx.idp_refresh_token + .clone() + .context("Missing idp_refresh_token")?, + ); + let http = default_oauth_http_client()?; + let token = EmaExchangeRequest::new(idp, resource_as, server_url, &refresh_token) + .with_scopes(manager.select_scopes(None, &[])) + .exchange(&http, &http) + .await?; + + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(server_url) + .auth_header(token.access_token.secret()), + ); + let client = BasicClientHandler + .serve_with_lifecycle(transport, conformance_lifecycle()) + .await?; + let tools = client.list_tools(Default::default()).await?; + for tool in tools.tools { + let args = build_tool_arguments(&tool); + client.call_tool(call_tool_params(tool.name, args)).await?; + } + client.cancel().await?; + Ok(()) +} + /// Cross-app access flow (SEP-1046 extension). async fn run_cross_app_access_client( server_url: &str, @@ -1110,6 +1189,11 @@ async fn run_scenario( "auth/client-credentials-basic" => run_client_credentials_basic(server_url, ctx).await?, "auth/client-credentials-jwt" => run_client_credentials_jwt(server_url, ctx).await?, + // Auth - enterprise-managed authorization with a refresh-token subject + "auth/enterprise-managed-authorization-refresh-token" => { + run_ema_refresh_token_client(server_url, ctx).await? + } + // Auth - cross-app access "auth/cross-app-access-complete-flow" => { run_cross_app_access_client(server_url, ctx).await? diff --git a/crates/rmcp-macros/CHANGELOG.md b/crates/rmcp-macros/CHANGELOG.md index 910257a28..67154ba86 100644 --- a/crates/rmcp-macros/CHANGELOG.md +++ b/crates/rmcp-macros/CHANGELOG.md @@ -7,6 +7,22 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [3.3.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-macros-v3.2.0...rmcp-macros-v3.3.0) - 2026-09-10 + +### Added + +- *(macros)* reject empty tool_router ([#1233](https://github.com/modelcontextprotocol/rust-sdk/pull/1233)) + +## [3.2.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-macros-v3.1.4...rmcp-macros-v3.2.0) - 2026-08-31 + +### Added + +- add request-state key rotation ([#1128](https://github.com/modelcontextprotocol/rust-sdk/pull/1128)) + +### Fixed + +- allow concurrent streamable http requests ([#1186](https://github.com/modelcontextprotocol/rust-sdk/pull/1186)) + ## [3.1.3](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-macros-v3.1.2...rmcp-macros-v3.1.3) - 2026-08-17 ### Fixed diff --git a/crates/rmcp-macros/src/lib.rs b/crates/rmcp-macros/src/lib.rs index 156e53b4a..176ead7cb 100644 --- a/crates/rmcp-macros/src/lib.rs +++ b/crates/rmcp-macros/src/lib.rs @@ -56,6 +56,7 @@ pub fn tool(attr: TokenStream, input: TokenStream) -> TokenStream { /// | `router` | `Ident` | The name of the router function to be generated. Defaults to `tool_router`. | /// | `vis` | `Visibility` | The visibility of the generated router function. Defaults to empty. | /// | `server_handler` | `flag` | When set, also emits `#[::rmcp::tool_handler]` on `impl ServerHandler for Self` so you can omit a separate `#[tool_handler]` block. | +/// | `allow_empty` | `flag` | When set, accepts an impl block with no `#[tool]` fn. Without it, an empty router is a compile error. | /// /// ## Example /// @@ -122,6 +123,37 @@ pub fn tool(attr: TokenStream, input: TokenStream) -> TokenStream { /// } /// } /// ``` +/// +/// ### Empty routers +/// +/// Collecting tools is this attribute's whole purpose, so an impl block with no `#[tool]` fn is a +/// compile error rather than a router that silently serves nothing. Pass `allow_empty` when that +/// is what you want: +/// +/// ```rust,ignore +/// #[tool_router(allow_empty)] +/// impl MyToolHandler {} +/// ``` +/// +/// The usual way to hit this by accident is a `macro_rules!` helper *inside* the impl block. An +/// attribute macro receives the unexpanded item, so `#[tool]` fns produced by such a helper are +/// invisible to `#[tool_router]`. Let the `macro_rules!` emit the whole annotated impl instead: +/// +/// ```rust,ignore +/// macro_rules! define_tools { +/// ($($name:ident => $description:literal),* $(,)?) => { +/// #[tool_router] +/// impl MyToolHandler { +/// $( +/// #[tool(description = $description)] +/// async fn $name(&self) -> String { stringify!($name).to_owned() } +/// )* +/// } +/// }; +/// } +/// +/// define_tools!(my_tool => "what my tool does"); +/// ``` #[proc_macro_attribute] pub fn tool_router(attr: TokenStream, input: TokenStream) -> TokenStream { tool_router::tool_router(attr.into(), input.into()) diff --git a/crates/rmcp-macros/src/tool_router.rs b/crates/rmcp-macros/src/tool_router.rs index edc8630b4..6efd6927f 100644 --- a/crates/rmcp-macros/src/tool_router.rs +++ b/crates/rmcp-macros/src/tool_router.rs @@ -17,6 +17,8 @@ pub struct ToolRouterAttribute { /// When set, also emit `#[::rmcp::tool_handler]` on `impl ServerHandler for Self` so callers /// can skip a separate `#[tool_handler]` block (expanded in a later macro pass). pub server_handler: bool, + /// When set, accept an impl block with no `#[tool]` fn instead of reporting an error. + pub allow_empty: bool, } impl Default for ToolRouterAttribute { @@ -25,6 +27,7 @@ impl Default for ToolRouterAttribute { router: format_ident!("tool_router"), vis: None, server_handler: false, + allow_empty: false, } } } @@ -35,6 +38,7 @@ pub fn tool_router(attr: TokenStream, input: TokenStream) -> syn::Result(input)?; // find all function marked with `#[rmcp::tool]` @@ -58,6 +62,19 @@ pub fn tool_router(attr: TokenStream, input: TokenStream) -> syn::Result syn::Result<()> { + let input = quote! { + impl Probe { + #[tool(description = "probe")] + async fn probe(&self) -> String { "probed".to_owned() } + } + }; + let generated = tool_router(TokenStream::new(), input)?.to_string(); + assert!(generated.contains("with_route"), "{generated}"); + Ok(()) + } + + #[test] + fn tool_router_allow_empty_generates_a_router_without_routes() -> syn::Result<()> { + let generated = tool_router(quote! { allow_empty }, quote! { impl Probe {} })?.to_string(); + assert!(generated.contains("fn tool_router"), "{generated}"); + assert!(!generated.contains("with_route"), "{generated}"); + Ok(()) + } } diff --git a/crates/rmcp/CHANGELOG.md b/crates/rmcp/CHANGELOG.md index 7bd421081..207358f54 100644 --- a/crates/rmcp/CHANGELOG.md +++ b/crates/rmcp/CHANGELOG.md @@ -9,8 +9,45 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added -- add `Peer::send_request_as` and option-aware typed request handles for - method-specific extension responses, plus an additive raw-response transport hook +- add `Peer::send_request_as` and option-aware typed request handles that + preserve method-specific extension response fields across supported built-in + transports + +### Fixed + +- reject malformed, duplicate, non-UTF-8, and disallowed HTTP Origin headers consistently +- compare HTTP origins using normalized effective ports instead of treating an omitted configured port as a wildcard + +## [3.3.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.2.0...rmcp-v3.3.0) - 2026-09-10 + +### Added + +- add ServerHandler::negotiate_initialize ([#1247](https://github.com/modelcontextprotocol/rust-sdk/pull/1247)) +- *(macros)* reject empty tool_router ([#1233](https://github.com/modelcontextprotocol/rust-sdk/pull/1233)) +- *(auth)* add enterprise refresh-token and ID-JAG exchanges ([#1234](https://github.com/modelcontextprotocol/rust-sdk/pull/1234)) + +### Fixed + +- *(sse)* saturate exponential reconnect backoff to avoid overflow panic ([#1231](https://github.com/modelcontextprotocol/rust-sdk/pull/1231)) +- resolve clippy warnings across workspace ([#1195](https://github.com/modelcontextprotocol/rust-sdk/pull/1195)) +- *(auth)* unify refresh checks and error handling ([#1236](https://github.com/modelcontextprotocol/rust-sdk/pull/1236)) + +### Other + +- *(deps)* update process-wrap requirement from 9.0 to 10.0 ([#1229](https://github.com/modelcontextprotocol/rust-sdk/pull/1229)) + +## [3.2.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.1.4...rmcp-v3.2.0) - 2026-08-31 + +### Added + +- *(auth)* coordinate OAuth refreshes through credential stores ([#1232](https://github.com/modelcontextprotocol/rust-sdk/pull/1232)) +- add request-state key rotation ([#1128](https://github.com/modelcontextprotocol/rust-sdk/pull/1128)) + +### Fixed + +- keep initialize on legacy protocol versions ([#1228](https://github.com/modelcontextprotocol/rust-sdk/pull/1228)) +- *(transport)* fall back after sessionless HTTP discover rejections ([#1211](https://github.com/modelcontextprotocol/rust-sdk/pull/1211)) +- allow concurrent streamable http requests ([#1186](https://github.com/modelcontextprotocol/rust-sdk/pull/1186)) ## [3.1.4](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.1.3...rmcp-v3.1.4) - 2026-08-18 diff --git a/crates/rmcp/Cargo.toml b/crates/rmcp/Cargo.toml index a4e2a42b4..8752c6ac9 100644 --- a/crates/rmcp/Cargo.toml +++ b/crates/rmcp/Cargo.toml @@ -18,6 +18,7 @@ exhaustive_enums = "warn" features = [ "auth", "auth-client-credentials-jwt", + "auth-enterprise-managed", "base64", "client", "client-side-sse", @@ -87,7 +88,7 @@ url = { version = "2.4", optional = true } tower-service = { version = "0.3", optional = true } # for child process transport -process-wrap = { version = "9.0", features = ["tokio1"], optional = true } +process-wrap = { version = "10.0", features = ["tokio1"], optional = true } # for cross-platform executable path resolution which = { version = "8", optional = true } @@ -196,10 +197,11 @@ transport-streamable-http-server-session = [ tower = ["dep:tower-service"] auth = ["dep:async-trait", "dep:oauth2", "__reqwest", "dep:url"] auth-client-credentials-jwt = ["auth", "dep:jsonwebtoken", "uuid"] +auth-enterprise-managed = ["auth", "base64"] schemars = ["dep:schemars"] [dev-dependencies] -tokio = { version = "1", features = ["full"] } +tokio = { version = "1", features = ["full", "test-util"] } schemars = { version = "1.1.0", features = ["chrono04"] } axum = { version = "0.8", default-features = false, features = ["http1", "tokio"] } hyper = { version = "1", features = ["server", "http1"] } diff --git a/crates/rmcp/README.md b/crates/rmcp/README.md index bb7837e84..f34469ef4 100644 --- a/crates/rmcp/README.md +++ b/crates/rmcp/README.md @@ -24,6 +24,7 @@ For **getting started**, **usage guides**, and **full MCP feature documentation* | `macros` | `#[tool]` / `#[prompt]` macros (re-exports [`rmcp-macros`](../rmcp-macros)) | ✅ | | `schemars` | JSON Schema generation for tool definitions | | | `auth` | OAuth 2.0 authentication support | | +| `auth-enterprise-managed` | EMA/XAA refresh-token and ID-JAG exchanges for registered public and confidential clients (includes `auth`) | | | `elicitation` | Elicitation support | | ### Transport features @@ -45,6 +46,10 @@ For **getting started**, **usage guides**, and **full MCP feature documentation* | `reqwest-native-tls` | Uses platform-native TLS (OpenSSL / Secure Transport / SChannel) | | `reqwest-tls-no-provider` | Uses rustls without a default crypto provider (bring your own) | +For enterprise-managed authorization, enable `auth-enterprise-managed` and a TLS +backend such as `reqwest`. See the [EMA/XAA guide](../../docs/OAUTH_SUPPORT.md#enterprise-managed-authorization-emaxaa) +for client authentication and an MCP connection example. + ## Transports The transport layer is pluggable. Two built-in pairs cover the most common cases: diff --git a/crates/rmcp/src/handler/server.rs b/crates/rmcp/src/handler/server.rs index 6dc7883ed..70985527c 100644 --- a/crates/rmcp/src/handler/server.rs +++ b/crates/rmcp/src/handler/server.rs @@ -321,13 +321,59 @@ macro_rules! server_handler_methods { context: RequestContext, ) -> impl Future> + MaybeSendFuture + '_ { context.peer.set_peer_info(request.clone()); + std::future::ready(self.negotiate_initialize(&request)) + } + /// Build the `initialize` response for `request`, negotiating the + /// protocol version against [`Self::supported_protocol_versions`]. + /// + /// This is the whole body of the default [`Self::initialize`] minus its + /// `set_peer_info` side effect, so a server that overrides `initialize` + /// to add its own can call this instead of restating the negotiation + /// rule: + /// + /// ``` + /// use rmcp::{ + /// ErrorData as McpError, RoleServer, ServerHandler, + /// model::{InitializeRequestParams, InitializeResult, ServerInfo}, + /// service::RequestContext, + /// }; + /// + /// struct MyServer; + /// + /// impl ServerHandler for MyServer { + /// fn get_info(&self) -> ServerInfo { + /// ServerInfo::default() + /// } + /// + /// async fn initialize( + /// &self, + /// request: InitializeRequestParams, + /// context: RequestContext, + /// ) -> Result { + /// // ... record telemetry, register the peer, etc. + /// context.peer.set_peer_info(request.clone()); + /// self.negotiate_initialize(&request) + /// } + /// } + /// ``` + /// + /// # Errors + /// + /// Returns [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`] when this server + /// supports no version that still has an `initialize` handshake. + /// + /// [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`]: crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION + fn negotiate_initialize( + &self, + request: &InitializeRequestParams, + ) -> Result { let mut info = self.get_info(); info.protocol_version = negotiate_protocol_version( &request.protocol_version, - info.protocol_version, + std::mem::take(&mut info.protocol_version), &self.supported_protocol_versions(), - ); - std::future::ready(Ok(info)) + )?; + Ok(info) } /// Return the protocol versions supported by this server. /// @@ -336,6 +382,10 @@ macro_rules! server_handler_methods { /// list is advertised by [`Self::discover`], bounds what `initialize` /// negotiation may agree to, and is what per-request versions are /// validated against. + /// + /// To support everything up to some ceiling, use + /// [`ProtocolVersion::known_up_to`] rather than filtering + /// [`ProtocolVersion::KNOWN_VERSIONS`] by hand. fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { Cow::Borrowed(ProtocolVersion::KNOWN_VERSIONS) } @@ -618,6 +668,13 @@ macro_rules! impl_server_handler_for_wrapper { (**self).initialize(request, context) } + fn negotiate_initialize( + &self, + request: &InitializeRequestParams, + ) -> Result { + (**self).negotiate_initialize(request) + } + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { (**self).supported_protocol_versions() } diff --git a/crates/rmcp/src/model.rs b/crates/rmcp/src/model.rs index 6a1409870..be911f7f9 100644 --- a/crates/rmcp/src/model.rs +++ b/crates/rmcp/src/model.rs @@ -177,7 +177,7 @@ impl ProtocolVersion { /// First protocol version that requires SEP-2243 standard HTTP headers. pub const STANDARD_HEADERS: Self = Self::V_2026_07_28; - /// All protocol versions known to this SDK. + /// All protocol versions known to this SDK, oldest first. pub const KNOWN_VERSIONS: &[Self] = &[ Self::V_2024_11_05, Self::V_2025_03_26, @@ -190,6 +190,43 @@ impl ProtocolVersion { pub fn as_str(&self) -> &str { &self.0 } + + /// The known versions up to and including `max`, oldest first. + /// + /// Servers that implement every revision up to some ceiling can return + /// this from `supported_protocol_versions` instead of filtering + /// [`Self::KNOWN_VERSIONS`] by hand. `max` itself need not be a known + /// version; the result is empty when it predates all of them. + /// + /// The result borrows from [`Self::KNOWN_VERSIONS`], so call it directly + /// in the method body — it needs no `static` and no `LazyLock`: + /// + /// ```rust,ignore + /// const MAX_SUPPORTED: ProtocolVersion = ProtocolVersion::V_2025_11_25; + /// + /// fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + /// Cow::Borrowed(ProtocolVersion::known_up_to(&MAX_SUPPORTED)) + /// } + /// ``` + /// + /// ``` + /// # use rmcp::model::ProtocolVersion; + /// assert_eq!( + /// ProtocolVersion::known_up_to(&ProtocolVersion::V_2025_06_18), + /// &[ + /// ProtocolVersion::V_2024_11_05, + /// ProtocolVersion::V_2025_03_26, + /// ProtocolVersion::V_2025_06_18, + /// ], + /// ); + /// ``` + pub fn known_up_to(max: &Self) -> &'static [Self] { + let count = Self::KNOWN_VERSIONS + .iter() + .take_while(|version| version.as_str() <= max.as_str()) + .count(); + &Self::KNOWN_VERSIONS[..count] + } } impl Serialize for ProtocolVersion { @@ -4643,6 +4680,51 @@ mod tests { use super::*; + #[test] + fn known_versions_are_ordered_oldest_first() { + // `known_up_to` walks the list as a sorted prefix. + assert!( + ProtocolVersion::KNOWN_VERSIONS + .windows(2) + .all(|pair| pair[0].as_str() < pair[1].as_str()) + ); + } + + #[test] + fn known_up_to_includes_the_ceiling_itself() { + assert_eq!( + ProtocolVersion::known_up_to(&ProtocolVersion::V_2024_11_05), + &[ProtocolVersion::V_2024_11_05] + ); + } + + #[test] + fn known_up_to_the_newest_version_yields_every_known_version() { + assert_eq!( + ProtocolVersion::known_up_to(&ProtocolVersion::V_2026_07_28), + ProtocolVersion::KNOWN_VERSIONS + ); + } + + #[test] + fn known_up_to_an_unknown_ceiling_stops_at_the_versions_below_it() { + let unknown = ProtocolVersion(Cow::Borrowed("2025-07-01")); + assert_eq!( + ProtocolVersion::known_up_to(&unknown), + &[ + ProtocolVersion::V_2024_11_05, + ProtocolVersion::V_2025_03_26, + ProtocolVersion::V_2025_06_18, + ] + ); + } + + #[test] + fn known_up_to_a_ceiling_below_every_known_version_is_empty() { + let ancient = ProtocolVersion(Cow::Borrowed("1999-01-01")); + assert!(ProtocolVersion::known_up_to(&ancient).is_empty()); + } + #[cfg(feature = "transport-streamable-http-client")] #[test] fn transport_closed_marker_accepts_only_the_process_local_token() { diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index c87612c20..32c59ff15 100644 --- a/crates/rmcp/src/service.rs +++ b/crates/rmcp/src/service.rs @@ -199,12 +199,17 @@ pub enum PeerRequestAssociation { Unknown { has_pending_outbound_request: bool }, } +/// Whether `version` predates `2026-07-28`, the revision that replaced the +/// `initialize` handshake with per-request metadata. +pub(crate) fn is_legacy_version(version: &ProtocolVersion) -> bool { + version.as_str() < ProtocolVersion::V_2026_07_28.as_str() +} + pub(crate) fn uses_legacy_lifecycle( protocol_version: Option<&ProtocolVersion>, uses_discover_lifecycle: bool, ) -> bool { - !uses_discover_lifecycle - && protocol_version.is_none_or(|version| version < &ProtocolVersion::V_2026_07_28) + !uses_discover_lifecycle && protocol_version.is_none_or(is_legacy_version) } pub(crate) fn peer_request_association( @@ -276,6 +281,12 @@ pub type RxJsonRpcMessage = JsonRpcMessage< ::PeerResp, ::PeerNot, >; +/// A received JSON-RPC message whose response result remains raw JSON. +/// +/// Requests and notifications retain the peer types defined by `R`; only a +/// successful response's result is represented as [`serde_json::Value`]. This +/// prevents extension fields from being lost to the role's core response union +/// before a typed request can deserialize them. pub type RawRxJsonRpcMessage = JsonRpcMessage<::PeerReq, serde_json::Value, ::PeerNot>; @@ -460,7 +471,7 @@ impl> DynService for S { } use std::{ - collections::{HashMap, VecDeque}, + collections::HashMap, ops::Deref, sync::{Arc, atomic::AtomicU64}, time::Duration, @@ -544,11 +555,12 @@ type ProgressTimeoutWatchers = Arc = (mpsc::Sender, usize); type SubscriptionChannelMap = HashMap>; -/// A handle to a remote request -/// -/// You can cancel it by call [`RequestHandle::cancel`] with a reason, +/// A handle to a remote request whose response resolves to `T`. /// -/// or wait for response by call [`RequestHandle::await_response`] +/// `T` defaults to the role's core peer-response union. Typed extension +/// requests created by [`Peer::send_request_as_with_option`] instead use the +/// caller's concrete result type. Call [`RequestHandle::cancel`] to cancel the +/// request with a reason, or [`RequestHandle::await_response`] to await it. #[derive(Debug)] #[non_exhaustive] pub struct RequestHandle::PeerResp> { @@ -748,6 +760,15 @@ impl PendingResponder { Self::Typed(responder) => responder(Ok(value)), } } + + fn send_standard_result(self, value: R::PeerResp) { + match self { + Self::Standard(responder) => { + let _ = responder.send(Ok(value)); + } + Self::Typed(responder) => responder(Err(ServiceError::RawResponseUnavailable)), + } + } } #[derive(Debug)] @@ -909,6 +930,14 @@ impl Peer { /// concrete type avoids ambiguous `#[serde(untagged)]` union matching. /// The caller is responsible for pairing the request method with its /// correct result type. + /// + /// # Errors + /// + /// Returns [`ServiceError::RawResponseUnavailable`] without sending the + /// request when the selected transport cannot preserve raw result JSON. + /// Returns [`ServiceError::ResponseDeserialization`] when the received + /// result does not deserialize as `T`. Transport, peer, timeout, and + /// cancellation errors are returned unchanged. pub async fn send_request_as(&self, request: R::Req) -> Result where T: serde::de::DeserializeOwned + Send + 'static, @@ -921,6 +950,9 @@ impl Peer { /// Send a typed request with the same timeout, metadata, progress, and /// cancellation lifecycle available to core requests. + /// + /// Raw-response support and error behavior are the same as + /// [`Peer::send_request_as`]. pub async fn send_request_as_with_option( &self, request: R::Req, @@ -1463,7 +1495,6 @@ where let current_span = tracing::Span::current(); let handle = spawn_service_task(async move { let mut transport = transport.into_transport(); - let mut batch_messages = VecDeque::>::new(); let mut send_task_set = tokio::task::JoinSet::::new(); let mut response_send_tasks = tokio::task::JoinSet::<()>::new(); #[derive(Debug)] @@ -1482,16 +1513,14 @@ where enum Event { ProxyMessage(PeerSinkMessage), PeerMessage(RawRxJsonRpcMessage), + LegacyPeerMessage(RxJsonRpcMessage), ToSink(TxJsonRpcMessage), SendTaskResult(SendTaskResult), ResponseSendTaskResult(Result<(), tokio::task::JoinError>), } let quit_reason = loop { - let evt = if let Some(m) = batch_messages.pop_front() { - Event::PeerMessage(m) - } else { - tokio::select! { + let evt = tokio::select! { m = sink_proxy_rx.recv(), if !sink_proxy_rx.is_closed() => { if let Some(m) = m { Event::ToSink(m) @@ -1499,9 +1528,15 @@ where continue } } - m = transport.receive_raw() => { - if let Some(m) = m { - Event::PeerMessage(m) + m = async { + if T::preserves_raw_responses() { + transport.receive_raw().await.map(Event::PeerMessage) + } else { + transport.receive().await.map(Event::LegacyPeerMessage) + } + } => { + if let Some(event) = m { + event } else { // input stream closed tracing::info!("input stream terminated"); @@ -1539,7 +1574,31 @@ where tracing::info!("task cancelled"); break QuitReason::Cancelled } + }; + + let evt = match evt { + Event::LegacyPeerMessage(JsonRpcMessage::Response(JsonRpcResponse { + result, + id, + .. + })) => { + if let Some(responder) = + remove_pending_request(&mut local_responder_pool, &id) + { + responder.send_standard_result(result); + } + continue; + } + Event::LegacyPeerMessage(JsonRpcMessage::Request(request)) => { + Event::PeerMessage(JsonRpcMessage::Request(request)) + } + Event::LegacyPeerMessage(JsonRpcMessage::Notification(notification)) => { + Event::PeerMessage(JsonRpcMessage::Notification(notification)) + } + Event::LegacyPeerMessage(JsonRpcMessage::Error(error)) => { + Event::PeerMessage(JsonRpcMessage::Error(error)) } + event => event, }; tracing::trace!(?evt, "new event"); @@ -1820,6 +1879,7 @@ where responder.send_error(service_error); } } + Event::LegacyPeerMessage(_) => unreachable!("legacy messages are normalized above"), } }; diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index 19719f593..51ad5445a 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -464,27 +464,62 @@ where } } -/// Echoes the client-requested version if the server supports it; otherwise -/// returns `server_fallback`. +/// Echoes the client-requested version if the server can serve it over the +/// `initialize` handshake; otherwise returns a legacy version the server does +/// support. /// /// `server_supported` comes from [`Service::supported_protocol_versions`], so a /// server that narrows that list is never made to answer `initialize` with a -/// version it cannot serve. +/// version it cannot serve. `2026-07-28` replaced the handshake with +/// per-request metadata, so a client naming that revision or later is answered +/// with the server's newest legacy version instead. +/// +/// # Errors +/// +/// Returns [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`] when the server supports +/// no version that still has an `initialize` handshake. +/// +/// [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`]: crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION pub(crate) fn negotiate_protocol_version( client_requested: &ProtocolVersion, server_fallback: ProtocolVersion, server_supported: &[ProtocolVersion], -) -> ProtocolVersion { - if server_supported.contains(client_requested) { - client_requested.clone() +) -> Result { + if is_legacy_version(client_requested) && server_supported.contains(client_requested) { + return Ok(client_requested.clone()); + } + let legacy_fallback = if is_legacy_version(&server_fallback) { + Some(server_fallback) } else { + newest_legacy_version(server_supported) + }; + let Some(legacy_fallback) = legacy_fallback else { tracing::warn!( client_requested = %client_requested, - server_fallback = %server_fallback, - "client requested unsupported protocol version; falling back to server default" + "server supports no protocol version with an initialize handshake; rejecting" ); - server_fallback - } + return Err(ErrorData::unsupported_protocol_version( + client_requested.clone(), + server_supported, + )); + }; + // Falling back is the designed answer for a pinned client, and stateless + // HTTP re-runs it on every request, so this is not a warning. + tracing::debug!( + client_requested = %client_requested, + server_fallback = %legacy_fallback, + "client requested a protocol version unavailable over initialize; falling back to server default" + ); + Ok(legacy_fallback) +} + +/// The newest of `versions` that still has an `initialize` handshake. +fn newest_legacy_version(versions: &[ProtocolVersion]) -> Option { + versions + .iter() + .filter(|version| is_legacy_version(version)) + .max_by(|left, right| left.as_str().cmp(right.as_str())) + .cloned() } fn missing_request_metadata_error(missing: &[&str]) -> ErrorData { @@ -497,6 +532,27 @@ fn missing_request_metadata_error(missing: &[&str]) -> ErrorData { ) } +/// Sends `error` as the response to the `initialize` request and reports it as +/// the reason the handshake failed. +async fn report_initialize_failure( + transport: &mut T, + error: ErrorData, + id: RequestId, +) -> ServerInitializeError +where + T: Transport + 'static, +{ + match transport + .send(ServerJsonRpcMessage::error(error.clone(), Some(id))) + .await + { + Ok(()) => ServerInitializeError::InitializeFailed(error), + Err(send_error) => { + ServerInitializeError::transport::(send_error, "sending error response") + } + } +} + async fn serve_server_with_ct_inner( service: S, transport: T, @@ -593,27 +649,21 @@ where peer: peer.clone(), }; // Send initialize response - let init_response = service.handle_request(request, context).await; - let mut init_response = match init_response { + let mut init_response = match service.handle_request(request, context).await { Ok(ServerResult::InitializeResult(init_response)) => init_response, Ok(result) => { return Err(ServerInitializeError::UnexpectedInitializeResponse(result)); } - Err(e) => { - transport - .send(ServerJsonRpcMessage::error(e.clone(), Some(id))) - .await - .map_err(|error| { - ServerInitializeError::transport::(error, "sending error response") - })?; - return Err(ServerInitializeError::InitializeFailed(e)); - } + Err(e) => return Err(report_initialize_failure(&mut transport, e, id).await), }; - init_response.protocol_version = negotiate_protocol_version( + init_response.protocol_version = match negotiate_protocol_version( &requested_protocol_version, init_response.protocol_version, &service.supported_protocol_versions(), - ); + ) { + Ok(version) => version, + Err(e) => return Err(report_initialize_failure(&mut transport, e, id).await), + }; // Update peer_info so context.protocol_version() reflects the negotiated // version in all subsequent request handlers. negotiated_peer_info.protocol_version = init_response.protocol_version.clone(); diff --git a/crates/rmcp/src/transport.rs b/crates/rmcp/src/transport.rs index 332f683f8..bf45d287b 100644 --- a/crates/rmcp/src/transport.rs +++ b/crates/rmcp/src/transport.rs @@ -100,7 +100,7 @@ pub use auth::JwtSigningAlgorithm; #[cfg(feature = "auth")] pub use auth::{ AuthClient, AuthError, AuthorizationManager, AuthorizationRequest, AuthorizationSession, - AuthorizedHttpClient, ClientCredentialsConfig, CredentialStore, + AuthorizedHttpClient, ClientCredentialsConfig, CredentialRefreshGuard, CredentialStore, EXTENSION_OAUTH_CLIENT_CREDENTIALS, InMemoryCredentialStore, InMemoryStateStore, OAuthHttpClient, OAuthHttpClientError, OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, ScopeUpgradeConfig, StateStore, StoredAuthorizationState, StoredCredentials, @@ -147,19 +147,32 @@ where /// Whether this transport preserves response result bodies as raw JSON. /// - /// Typed extension requests require this capability. Existing transports - /// default to `false` because decoding into the role response union can - /// irreversibly discard extension fields. + /// Typed extension requests require this capability. Returning `true` is a + /// contract that [`Self::receive_raw`] returns the original JSON result for + /// every response path without first decoding it through the role response + /// union. A transport that cannot satisfy that contract must retain the + /// default `false`; typed requests then fail with + /// [`crate::service::ServiceError::RawResponseUnavailable`] before they are + /// sent. + /// + /// The built-in async-read/write, child-process, reqwest HTTP, Unix-socket + /// HTTP transports preserve raw responses. Authenticated HTTP wrappers + /// forward the wrapped backend's capability, and a `WorkerTransport` + /// forwards the capability declared by its worker. + /// Existing custom transports default to `false` for source and behavior + /// compatibility. fn preserves_raw_responses() -> bool { false } /// Receive a message while preserving the raw JSON-RPC result value. /// - /// Transports should override this method when they can retain the raw - /// response body. The default preserves compatibility for existing custom - /// transports, but an extension result already decoded through the role's - /// response union cannot recover information discarded by that union. + /// A transport that returns `true` from [`Self::preserves_raw_responses`] + /// must override this method and retain raw results on every response path. + /// The default adapts the existing typed receive API for custom transports; + /// it does not make typed extension requests available because an extension + /// result already decoded through the role response union cannot recover + /// information discarded by that union. fn receive_raw(&mut self) -> impl Future>> + Send { async move { self.receive().await.map(|message| match message { diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index 2eae2b220..759044b0c 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -28,6 +28,9 @@ use tracing::{debug, warn}; use crate::transport::common::http_header::HEADER_MCP_PROTOCOL_VERSION; +#[cfg(feature = "auth-enterprise-managed")] +pub mod enterprise; + const DEFAULT_HTTP_TIMEOUT: Duration = Duration::from_secs(30); const MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES: usize = 1024 * 1024; const MAX_OAUTH_DISCOVERY_REDIRECTS: usize = 10; @@ -99,6 +102,19 @@ pub trait OAuthHttpClient: Send + Sync { fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_>; } +/// Create an OAuth HTTP client with the SDK's default reqwest configuration. +/// +/// Honors each request's redirect policy, with a 30-second timeout and bounded +/// response bodies. Enable a TLS feature such as `reqwest` for HTTPS requests. +/// Implement [`OAuthHttpClient`] instead when custom network policy is required. +pub fn default_oauth_http_client() -> Result { + let client = ReqwestClient::builder() + .timeout(DEFAULT_HTTP_TIMEOUT) + .build() + .map_err(|error| AuthError::InternalError(error.to_string()))?; + ReqwestOAuthHttpClient::new(client) +} + struct ReqwestOAuthHttpClient { follow_redirects: ReqwestClient, stop_redirects: ReqwestClient, @@ -255,11 +271,32 @@ impl StoredCredentials { } } +/// An owned guard held across a credential refresh and its save. +/// +/// Stores can wrap a file lock, an owned mutex guard, or another coordination +/// primitive. Dropping this value releases the guard. +#[must_use = "dropping the guard releases refresh coordination"] +pub struct CredentialRefreshGuard { + _guard: Box, +} + +impl CredentialRefreshGuard { + /// Wrap a guard whose lifetime coordinates access to the stored credentials. + pub fn new(guard: impl Send + 'static) -> Self { + Self { + _guard: Box::new(guard), + } + } +} + /// Trait for storing and retrieving OAuth2 credentials /// /// Implementations of this trait can provide custom storage backends /// for OAuth2 credentials, such as file-based storage, keychain integration, /// or database storage. +/// +/// Return [`AuthError::CredentialStoreError`] for backend or locking failures +/// so they remain distinct from errors requiring reauthorization. #[async_trait] pub trait CredentialStore: Send + Sync { async fn load(&self) -> Result, AuthError>; @@ -267,6 +304,16 @@ pub trait CredentialStore: Send + Sync { async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError>; async fn clear(&self) -> Result<(), AuthError>; + + /// Optionally coordinate refreshes that share these credentials. + /// + /// The manager acquires this guard before loading credentials and retains it + /// through the token request and save. `load` and `save` must not reacquire + /// the same lock. Writers that bypass the guard are not coordinated with it. + /// The default does not coordinate refreshes. + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + Ok(None) + } } /// In-memory credential store (default implementation) @@ -515,6 +562,9 @@ pub enum AuthError { #[error("OAuth refresh token was rejected: {0}")] TokenRefreshRejected(String), + #[error("OAuth credential store failed: {0}")] + CredentialStoreError(String), + #[error("HTTP error: {0}")] HttpError(#[from] reqwest::Error), @@ -1274,13 +1324,9 @@ impl AuthorizationManager { /// create new auth manager with base url pub async fn new(base_url: U) -> Result { - let http_client = ReqwestClient::builder() - .timeout(DEFAULT_HTTP_TIMEOUT) - .build() - .map_err(|e| AuthError::InternalError(e.to_string()))?; Self::new_inner( base_url, - Arc::new(ReqwestOAuthHttpClient::new(http_client)?), + Arc::new(default_oauth_http_client()?), OAuthHttpRedirectPolicy::Stop, ) .await @@ -2203,8 +2249,21 @@ impl AuthorizationManager { .as_ref() .ok_or_else(|| AuthError::InternalError("OAuth client not configured".to_string()))?; + // Held for the rest of this function so the load, the exchange, and the + // save stay inside one guarded section. + let _refresh_guard = self.credential_store.acquire_refresh_guard().await?; let stored = self.credential_store.load().await?; let stored_credentials = stored.ok_or(AuthError::AuthorizationRequired)?; + // Refreshing with another client's stored token would put that token on a + // request authenticated as this client. + if stored_credentials.client_id != oauth_client.client_id().as_str() { + tracing::warn!( + stored_client_id = stored_credentials.client_id.as_str(), + configured_client_id = oauth_client.client_id().as_str(), + "stored credentials belong to a different client; reauthorization required" + ); + return Err(AuthError::AuthorizationRequired); + } let current_credentials = stored_credentials .token_response .ok_or(AuthError::AuthorizationRequired)?; @@ -2221,6 +2280,7 @@ impl AuthorizationManager { .add_extra_param("resource", self.oauth_resource().await); let mut refresh_scopes = stored_credentials.granted_scopes; self.add_offline_access_if_supported(&mut refresh_scopes); + let requested_scopes = refresh_scopes.clone(); for scope in refresh_scopes { refresh_request = refresh_request.add_scope(Scope::new(scope)); } @@ -2246,9 +2306,12 @@ impl AuthorizationManager { token_result.set_refresh_token(Some(refresh_token_value)); } - let granted_scopes: Vec = match token_result.scopes() { - Some(scopes) => scopes.iter().map(|s| s.to_string()).collect(), - None => self.current_scopes.read().await.clone(), + let response_scopes = token_result + .scopes() + .map(|scopes| scopes.iter().map(|s| s.to_string()).collect()); + let granted_scopes = { + let current = self.current_scopes.read().await; + Self::resolve_granted_scopes(response_scopes, &requested_scopes, ¤t) }; *self.current_scopes.write().await = granted_scopes.clone(); @@ -3904,17 +3967,19 @@ mod tests { sync::{Arc, Mutex as StdMutex}, }; - use oauth2::{AuthType, CsrfToken, HttpResponse, PkceCodeVerifier}; + use oauth2::{AuthType, CsrfToken, HttpResponse, PkceCodeVerifier, TokenResponse}; use reqwest::StatusCode; use rstest::rstest; + use tokio::sync::{Mutex, OwnedMutexGuard, Semaphore}; use url::Url; use super::{ AuthError, AuthorizationCallback, AuthorizationManager, AuthorizationMetadata, - AuthorizationMetadataSource, AuthorizationRequest, AuthorizationSession, CredentialStore, - InMemoryCredentialStore, InMemoryStateStore, OAuthClientConfig, OAuthHttpClient, - OAuthHttpClientError, OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, - ScopeUpgradeConfig, StateStore, StoredAuthorizationState, is_https_url, + AuthorizationMetadataSource, AuthorizationRequest, AuthorizationSession, + CredentialRefreshGuard, CredentialStore, InMemoryCredentialStore, InMemoryStateStore, + OAuthClientConfig, OAuthHttpClient, OAuthHttpClientError, OAuthHttpClientFuture, + OAuthHttpRedirectPolicy, OAuthHttpRequest, ScopeUpgradeConfig, StateStore, + StoredAuthorizationState, is_https_url, }; use crate::transport::auth::VendorExtraTokenFields; @@ -4000,6 +4065,58 @@ mod tests { ); } + #[tokio::test] + async fn default_oauth_http_client_honors_redirect_policy() { + use axum::{Router, routing::post}; + + let received = Arc::new(StdMutex::new(Vec::new())); + let capture = Arc::clone(&received); + let app = Router::new() + .route( + "/redirect", + post(|| async { (StatusCode::TEMPORARY_REDIRECT, [("location", "/token")]) }), + ) + .route( + "/token", + post(move |body: String| { + let capture = Arc::clone(&capture); + async move { + capture.lock().unwrap().push(body); + StatusCode::OK + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let endpoint = format!("http://{}/redirect", listener.local_addr().unwrap()); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let client = super::default_oauth_http_client().unwrap(); + + for (policy, expected) in [ + ( + OAuthHttpRedirectPolicy::Stop, + StatusCode::TEMPORARY_REDIRECT, + ), + (OAuthHttpRedirectPolicy::Follow, StatusCode::OK), + ] { + let request = oauth2::http::Request::builder() + .method("POST") + .uri(&endpoint) + .body(b"credential-sentinel".to_vec()) + .unwrap(); + let response = client + .execute(OAuthHttpRequest::new(request, policy)) + .await + .unwrap(); + assert_eq!(response.status(), expected); + let expected_bodies = if policy == OAuthHttpRedirectPolicy::Stop { + vec![] + } else { + vec!["credential-sentinel".to_owned()] + }; + assert_eq!(*received.lock().unwrap(), expected_bodies); + } + } + #[tokio::test] async fn default_http_client_preserves_connection_failure_cause() { let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); @@ -8342,4 +8459,349 @@ mod tests { "a rotated refresh token from the response should replace the old one" ); } + + #[derive(Clone)] + struct RefreshStore { + credentials: InMemoryCredentialStore, + lock: Arc>, + events: Arc>>, + guard_requested: Arc, + save_started: Arc, + save_gate: Option>, + fail_at: Option<&'static str>, + } + + struct ObservedRefreshGuard { + _lock: OwnedMutexGuard<()>, + events: Arc>>, + } + + impl Drop for ObservedRefreshGuard { + fn drop(&mut self) { + self.events.lock().unwrap().push("release"); + } + } + + #[async_trait::async_trait] + impl CredentialStore for RefreshStore { + async fn load(&self) -> Result, AuthError> { + self.events.lock().unwrap().push("load"); + if self.fail_at == Some("load") { + return Err(AuthError::CredentialStoreError("load failed".into())); + } + self.credentials.load().await + } + + async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError> { + self.events.lock().unwrap().push("save"); + self.save_started.add_permits(1); + if let Some(gate) = &self.save_gate { + gate.acquire().await.unwrap().forget(); + } + if self.fail_at == Some("save") { + return Err(AuthError::CredentialStoreError("save failed".into())); + } + self.credentials.save(credentials).await?; + self.events.lock().unwrap().push("saved"); + Ok(()) + } + + async fn clear(&self) -> Result<(), AuthError> { + self.credentials.clear().await + } + + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + self.events.lock().unwrap().push("acquire"); + self.guard_requested.add_permits(1); + if self.fail_at == Some("guard") { + return Err(AuthError::CredentialStoreError("guard failed".into())); + } + let lock = self.lock.clone().lock_owned().await; + self.events.lock().unwrap().push("acquired"); + Ok(Some(CredentialRefreshGuard::new(ObservedRefreshGuard { + _lock: lock, + events: self.events.clone(), + }))) + } + } + + async fn refresh_store() -> RefreshStore { + let credentials = StoredCredentials::new( + "my-client".into(), + Some(make_token_response_with_refresh("old-token", "old-refresh")), + vec!["read".into()], + Some(AuthorizationManager::now_epoch_secs()), + ); + let credential_store = InMemoryCredentialStore::new(); + credential_store.save(credentials).await.unwrap(); + RefreshStore { + credentials: credential_store, + lock: Arc::new(Mutex::new(())), + events: Arc::new(StdMutex::new(Vec::new())), + guard_requested: Arc::new(Semaphore::new(0)), + save_started: Arc::new(Semaphore::new(0)), + save_gate: None, + fail_at: None, + } + } + + struct RefreshHttpClient { + recording: RecordingOAuthHttpClient, + events: Arc>>, + } + + impl OAuthHttpClient for RefreshHttpClient { + fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + self.events.lock().unwrap().push("provider"); + self.recording.execute(request) + } + } + + fn refresh_http_client(store: &RefreshStore) -> Arc { + Arc::new(RefreshHttpClient { + recording: RecordingOAuthHttpClient::with_responses( + [ + ("new-token", "new-refresh"), + ("latest-token", "latest-refresh"), + ] + .into_iter() + .map(|(access, refresh)| { + http_response( + 200, + serde_json::json!({ + "access_token": access, "token_type": "Bearer", + "expires_in": 3600, "refresh_token": refresh + }), + ) + }) + .collect(), + ), + events: store.events.clone(), + }) + } + + async fn refresh_manager( + store: RefreshStore, + http_client: Arc, + ) -> AuthorizationManager { + let mut manager = AuthorizationManager::new_with_oauth_http_client( + "https://mcp.example.com/mcp", + http_client, + ) + .await + .unwrap(); + manager.set_metadata(AuthorizationMetadata { + authorization_endpoint: "https://auth.example.com/authorize".into(), + token_endpoint: "https://auth.example.com/token".into(), + ..Default::default() + }); + manager.configure_client(test_client_config()).unwrap(); + manager.set_credential_store(store); + *manager.current_scopes.write().await = vec!["cached".into()]; + manager + } + + async fn wait_for_permits(semaphore: &Semaphore, count: u32) { + tokio::time::timeout( + std::time::Duration::from_secs(5), + semaphore.acquire_many(count), + ) + .await + .unwrap() + .unwrap() + .forget(); + } + + #[tokio::test] + async fn refresh_guard_spans_load_exchange_and_completed_save() { + let store = refresh_store().await; + let manager = refresh_manager(store.clone(), refresh_http_client(&store)).await; + + manager.refresh_token().await.unwrap(); + + assert_eq!( + *store.events.lock().unwrap(), + [ + "acquire", "acquired", "load", "provider", "save", "saved", "release" + ] + ); + let saved = store.credentials.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + "new-refresh" + ); + assert_eq!(saved.granted_scopes, ["read"]); + assert!(store.lock.try_lock().is_ok()); + } + + #[tokio::test] + async fn concurrent_refreshes_wait_for_save_and_use_the_latest_token() { + let mut store = refresh_store().await; + let save_gate = Arc::new(Semaphore::new(0)); + store.save_gate = Some(save_gate.clone()); + let http_client = refresh_http_client(&store); + let first_manager = refresh_manager(store.clone(), http_client.clone()).await; + let second_manager = refresh_manager(store.clone(), http_client.clone()).await; + + let first = tokio::spawn(async move { first_manager.refresh_token().await }); + wait_for_permits(&store.save_started, 1).await; + let second = tokio::spawn(async move { second_manager.refresh_token().await }); + wait_for_permits(&store.guard_requested, 2).await; + assert_eq!(http_client.recording.requests().len(), 1); + assert!(store.lock.try_lock().is_err()); + + save_gate.add_permits(2); + let (first, second) = tokio::join!(first, second); + assert_eq!(first.unwrap().unwrap().access_token().secret(), "new-token"); + assert_eq!( + second.unwrap().unwrap().access_token().secret(), + "latest-token" + ); + let refresh_tokens: Vec = http_client + .recording + .requests() + .iter() + .map(|request| { + url::form_urlencoded::parse(&request.body) + .find_map(|(key, value)| (key == "refresh_token").then(|| value.into_owned())) + .unwrap() + }) + .collect(); + assert_eq!(refresh_tokens, ["old-refresh", "new-refresh"]); + let saved = store.credentials.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + "latest-refresh" + ); + assert!(store.lock.try_lock().is_ok()); + } + + #[tokio::test] + async fn refresh_rejects_credentials_for_another_client() { + let store = refresh_store().await; + let mut credentials = store.credentials.load().await.unwrap().unwrap(); + credentials.client_id = "other-client".into(); + store.credentials.save(credentials).await.unwrap(); + let http_client = refresh_http_client(&store); + let manager = refresh_manager(store.clone(), http_client.clone()).await; + + assert!(matches!( + manager.refresh_token().await, + Err(AuthError::AuthorizationRequired) + )); + assert!(http_client.recording.requests().is_empty()); + assert!(store.lock.try_lock().is_ok()); + } + + #[tokio::test] + async fn refresh_rejects_credentials_for_another_client_without_a_guard() { + let (base_url, captured) = start_token_server().await; + let mut manager = manager_with_metadata(Some(AuthorizationMetadata { + authorization_endpoint: format!("{base_url}/authorize"), + token_endpoint: format!("{base_url}/token"), + ..Default::default() + })) + .await; + manager.configure_client(test_client_config()).unwrap(); + manager + .credential_store + .save(StoredCredentials::new( + "other-client".into(), + Some(make_token_response_with_refresh("old-token", "old-refresh")), + vec!["read".into()], + Some(AuthorizationManager::now_epoch_secs()), + )) + .await + .unwrap(); + + let error = manager.refresh_token().await.unwrap_err(); + + assert!( + matches!(error, AuthError::AuthorizationRequired), + "a client mismatch must require reauthorization, got: {error:?}" + ); + assert!( + captured.lock().unwrap().is_none(), + "a client mismatch must be caught before the refresh token leaves the process" + ); + } + + #[tokio::test] + async fn refresh_without_a_guard_keeps_stored_scopes_when_response_omits_them() { + // start_token_server answers without a `scope`, matching a provider that + // grants the request in full. + let (base_url, _captured) = start_token_server().await; + let mut manager = manager_with_metadata(Some(AuthorizationMetadata { + authorization_endpoint: format!("{base_url}/authorize"), + token_endpoint: format!("{base_url}/token"), + ..Default::default() + })) + .await; + manager.configure_client(test_client_config()).unwrap(); + manager + .credential_store + .save(StoredCredentials::new( + "my-client".into(), + Some(make_token_response_with_refresh("old-token", "old-refresh")), + vec!["read".into()], + Some(AuthorizationManager::now_epoch_secs()), + )) + .await + .unwrap(); + *manager.current_scopes.write().await = vec!["stale".into()]; + + manager.refresh_token().await.unwrap(); + + let saved = manager.credential_store.load().await.unwrap().unwrap(); + assert_eq!( + saved.granted_scopes, + ["read"], + "the stored grant outranks the per-process scope cache" + ); + assert_eq!( + manager.get_current_scopes().await, + ["read"], + "the refreshed grant must replace the stale scope cache" + ); + } + + #[rstest] + #[case("guard", 0)] + #[case("load", 0)] + #[case("save", 1)] + #[tokio::test] + async fn refresh_store_failures_release_the_guard( + #[case] phase: &'static str, + #[case] provider_requests: usize, + ) { + let mut store = refresh_store().await; + store.fail_at = Some(phase); + let http_client = refresh_http_client(&store); + let manager = refresh_manager(store.clone(), http_client.clone()).await; + + assert!(matches!(manager.refresh_token().await, + Err(AuthError::CredentialStoreError(message)) if message == format!("{phase} failed"))); + assert_eq!(http_client.recording.requests().len(), provider_requests); + let saved = store.credentials.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + "old-refresh" + ); + assert!(store.lock.try_lock().is_ok()); + } } diff --git a/crates/rmcp/src/transport/auth/enterprise.rs b/crates/rmcp/src/transport/auth/enterprise.rs new file mode 100644 index 000000000..226c8d531 --- /dev/null +++ b/crates/rmcp/src/transport/auth/enterprise.rs @@ -0,0 +1,767 @@ +//! Non-interactive enterprise-managed authorization (EMA/XAA) token exchanges. +//! +//! Exchanges an enterprise refresh token for an ID-JAG, then for an MCP access +//! token. Callers must discover and approve both servers and their registrations +//! before supplying a credential. This module does not discover servers, log in, +//! persist credentials, or decide when to reauthenticate. +//! +//! Configure each pre-registered client's approved authentication method separately: +//! HTTP Basic, a client secret in the request body, or a freshly signed JWT client +//! assertion. Public-client authentication requires the server's explicit approval. +//! This helper requires one MCP resource and does not implement Rich +//! Authorization Requests or DPoP. Redemption consumes the SDK's ID-JAG handle, +//! without automatic retries; server-side replay policy remains the server's responsibility. +//! +//! ID-JAG checks below enforce structure and claim bindings, not cryptographic +//! signature verification. Assertions come directly from the trusted IdP token +//! endpoint; the resource authorization server must verify their signatures. +//! +//! ```no_run +//! use oauth2::{ClientSecret, RefreshToken}; +//! use rmcp::transport::auth::{default_oauth_http_client, enterprise::*}; +//! +//! # async fn authorize(refresh: &RefreshToken, idp_secret: ClientSecret, resource_secret: ClientSecret) -> Result<(), Box> { +//! // Enable `auth-enterprise-managed` and a TLS feature such as `reqwest`. +//! let http = default_oauth_http_client()?; +//! let token = EmaExchangeRequest::new( +//! EmaAuthorizationServer::new("https://idp.example", "https://idp.example/token", "idp-client") +//! .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(idp_secret)), +//! EmaAuthorizationServer::new("https://as.example", "https://as.example/token", "mcp-client") +//! .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(resource_secret)), +//! "https://mcp.example", refresh, +//! ).with_scopes(["files.read"]).exchange(&http, &http).await?; +//! // Use token.access_token only for the approved MCP resource; never log it. +//! // token.scopes contains the final granted scopes, which may be narrower. +//! # Ok(()) } +//! ``` + +use std::{ + collections::HashSet, + sync::Arc, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use base64::{ + Engine, + engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}, +}; +use oauth2::{AccessToken, ClientSecret, RefreshToken}; +use serde::{Deserialize, de::DeserializeOwned}; +use thiserror::Error; +use url::{Host, Url}; + +use super::{ + DEFAULT_HTTP_TIMEOUT, MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES, OAuthHttpClient, + OAuthHttpRedirectPolicy, OAuthHttpRequest, +}; + +const ID_JAG_TOKEN_TYPE: &str = "urn:ietf:params:oauth:token-type:id-jag"; + +/// Authentication approved for a pre-registered client at one authorization server. +/// +/// The selected method is used as configured, without negotiation or fallback. +#[derive(Clone)] +#[non_exhaustive] +pub enum EmaClientAuthentication { + /// Public client (`token_endpoint_auth_method=none`), only if the server permits it. + None, + /// `client_secret_basic`, with OAuth form encoding before HTTP Basic encoding. + ClientSecretBasic(ClientSecret), + /// `client_secret_post`, for servers requiring credentials in the request body. + ClientSecretPost(ClientSecret), + /// Fresh JWT client assertions, such as `private_key_jwt` or `client_secret_jwt`. + /// Signing, key custody, claims, and the registered algorithm belong to the provider. + JwtAssertion(Arc), +} + +impl std::fmt::Debug for EmaClientAuthentication { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(match self { + Self::None => "None", + Self::ClientSecretBasic(_) => "ClientSecretBasic { .. }", + Self::ClientSecretPost(_) => "ClientSecretPost { .. }", + Self::JwtAssertion(_) => "JwtAssertion { .. }", + }) + } +} + +impl EmaClientAuthentication { + fn validate(&self) -> Result<(), EmaError> { + if let Self::ClientSecretBasic(secret) | Self::ClientSecretPost(secret) = self + && secret.secret().trim().is_empty() + { + return Err(EmaError::InvalidRequest("client secret must not be empty")); + } + Ok(()) + } +} + +/// A signed JWT used to authenticate the client, distinct from the ID-JAG grant. +pub struct EmaClientAssertion(String); + +impl EmaClientAssertion { + /// Wrap a fresh assertion without exposing it through `Debug`. + pub fn new(assertion: String) -> Self { + Self(assertion) + } +} + +impl std::fmt::Debug for EmaClientAssertion { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaClientAssertion { .. }") + } +} + +/// Creates client assertions on demand, allowing keys to remain in an external signer. +/// +/// Called once immediately before each token request, including delayed ID-JAG +/// redemption. Set `iss` and `sub` to the registered client identifier and `aud` +/// to the server's approved audience, with a short expiration and a fresh `jti`. +/// The SDK does not sign or validate these assertions. Provider failures are +/// sanitized and the call shares the token request's timeout. Cancellation may +/// occur when that deadline expires. +#[async_trait::async_trait] +pub trait EmaClientAssertionProvider: Send + Sync { + async fn create_assertion( + &self, + server: &EmaAuthorizationServer, + ) -> Result>; +} + +/// A trusted server with a pre-registered client and its approved authentication method. +#[derive(Clone)] +#[non_exhaustive] +pub struct EmaAuthorizationServer { + /// Exact issuer identifier from approved metadata. + pub issuer: String, + /// Token endpoint from that metadata. + pub token_endpoint: String, + /// Pre-registered client identifier. + pub client_id: String, + client_authentication: EmaClientAuthentication, +} + +impl EmaAuthorizationServer { + /// Use approved metadata and a public-client registration (`token_endpoint_auth_method=none`). + /// Set [`Self::with_client_authentication`] for a confidential client. + pub fn new( + issuer: impl Into, + token_endpoint: impl Into, + client_id: impl Into, + ) -> Self { + Self { + issuer: issuer.into(), + token_endpoint: token_endpoint.into(), + client_id: client_id.into(), + client_authentication: EmaClientAuthentication::None, + } + } + + /// Select the authentication method approved for this server's client registration. + pub fn with_client_authentication(mut self, authentication: EmaClientAuthentication) -> Self { + self.client_authentication = authentication; + self + } +} + +/// A refresh-token exchange bound to one MCP resource and two registered clients. +pub struct EmaExchangeRequest<'a> { + idp: EmaAuthorizationServer, + resource_as: EmaAuthorizationServer, + resource: &'a str, + refresh_token: &'a RefreshToken, + scopes: Vec, +} + +impl<'a> EmaExchangeRequest<'a> { + /// No scope parameter is sent until [`Self::with_scopes`] is used. + pub fn new( + idp: EmaAuthorizationServer, + resource_as: EmaAuthorizationServer, + resource: &'a str, + refresh_token: &'a RefreshToken, + ) -> Self { + Self { + idp, + resource_as, + resource, + refresh_token, + scopes: Vec::new(), + } + } + + /// Request distinct non-empty scope tokens; the IdP may narrow them. + /// An empty iterator omits `scope`, rather than requesting an empty grant. + pub fn with_scopes(mut self, scopes: I) -> Self + where + I: IntoIterator, + S: Into, + { + self.scopes = scopes.into_iter().map(Into::into).collect(); + self + } + + /// Obtain an ID-JAG without redirects or retries, checking its resource/client bindings. + /// The returned assertion does not retain the enterprise refresh token. + pub async fn exchange_id_jag( + self, + idp_http: &dyn OAuthHttpClient, + ) -> Result { + for endpoint in [ + self.resource, + &self.idp.issuer, + &self.idp.token_endpoint, + &self.resource_as.issuer, + &self.resource_as.token_endpoint, + ] { + validate_endpoint(endpoint)?; + } + if self.idp.issuer == self.resource_as.issuer { + return Err(EmaError::InvalidRequest( + "IdP and resource AS issuers must differ", + )); + } + if self.idp.client_id.trim().is_empty() + || self.resource_as.client_id.trim().is_empty() + || self.refresh_token.secret().trim().is_empty() + { + return Err(EmaError::InvalidRequest( + "client IDs and refresh token must not be empty", + )); + } + self.idp.client_authentication.validate()?; + self.resource_as.client_authentication.validate()?; + let requested: HashSet<&str> = self.scopes.iter().map(String::as_str).collect(); + if requested.len() != self.scopes.len() || self.scopes.iter().any(|s| !is_scope_token(s)) { + return Err(EmaError::InvalidRequest( + "scopes must be distinct non-empty tokens", + )); + } + let mut params = vec![ + ( + "grant_type", + "urn:ietf:params:oauth:grant-type:token-exchange", + ), + ("requested_token_type", ID_JAG_TOKEN_TYPE), + ( + "subject_token_type", + "urn:ietf:params:oauth:token-type:refresh_token", + ), + ("subject_token", self.refresh_token.secret()), + ("audience", self.resource_as.issuer.as_str()), + ("resource", self.resource), + ]; + let scope = self.scopes.join(" "); + if !scope.is_empty() { + params.push(("scope", &scope)); + } + let jag: IdJagResponse = post_form( + idp_http, + &self.idp, + ¶ms, + EmaExchangeStage::IdentityProvider, + None, + unix_time, + ) + .await?; + let (granted, expires_at) = jag.validate(&self, &requested)?; + Ok(EmaIdJag { + assertion: AccessToken::new(jag.access_token), + scopes: granted, + resource_as: self.resource_as, + resource: self.resource.to_owned(), + expires_at, + }) + } + + /// Perform both exchanges with independently routed HTTP clients and no automatic retries. + pub async fn exchange( + self, + idp_http: &dyn OAuthHttpClient, + resource_http: &dyn OAuthHttpClient, + ) -> Result { + self.exchange_id_jag(idp_http) + .await? + .exchange(resource_http) + .await + } +} + +/// An IdP-issued ID-JAG whose structure and bindings have been checked, not its signature. +pub struct EmaIdJag { + assertion: AccessToken, + scopes: HashSet, + resource_as: EmaAuthorizationServer, + resource: String, + expires_at: u64, +} + +impl std::fmt::Debug for EmaAuthorizationServer { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaAuthorizationServer { .. }") + } +} + +impl std::fmt::Debug for EmaExchangeRequest<'_> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaExchangeRequest { .. }") + } +} + +impl std::fmt::Debug for EmaIdJag { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaIdJag { .. }") + } +} + +impl EmaIdJag { + /// The assertion to present to the approved resource authorization server. + pub fn assertion(&self) -> &AccessToken { + &self.assertion + } + + /// The scope tokens carried by the assertion; an empty set means scope was omitted. + pub fn scopes(&self) -> &HashSet { + &self.scopes + } + + /// Redeem this assertion once, at its approved resource AS, without redirects or retries. + /// Consume the grant so this helper cannot accidentally replay it after a failed exchange. + pub async fn exchange(self, http: &dyn OAuthHttpClient) -> Result { + self.exchange_with_clock(http, unix_time).await + } + + async fn exchange_with_clock( + self, + http: &dyn OAuthHttpClient, + now: impl Fn() -> Result + Sync, + ) -> Result { + if self.expires_at <= now()? { + return Err(EmaError::InvalidRequest("ID-JAG expired before redemption")); + } + // Only the assertion carries authority: repeating resource/scope could undo narrowing. + let token: ResourceTokenResponse = post_form( + http, + &self.resource_as, + &[ + ("grant_type", "urn:ietf:params:oauth:grant-type:jwt-bearer"), + ("assertion", self.assertion.secret()), + ], + EmaExchangeStage::ResourceAuthorizationServer, + Some(self.expires_at), + now, + ) + .await?; + token.validate(&self.resource, self.scopes) + } +} + +/// A resource-bound bearer with secret-safe diagnostics and its optional lifetime. +#[derive(Clone)] +#[non_exhaustive] +pub struct EmaAccessToken { + pub access_token: AccessToken, + pub expires_in: Option, + /// Resource-AS scopes, or the ID-JAG scopes when the response omits `scope`. + /// Empty when both the ID-JAG and response omit scopes. + pub scopes: HashSet, +} + +impl std::fmt::Debug for EmaAccessToken { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaAccessToken { .. }") + } +} + +/// The endpoint that failed, allowing the caller to apply its own credential lifecycle policy. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub enum EmaExchangeStage { + IdentityProvider, + ResourceAuthorizationServer, +} + +/// Sanitized failures. Raw HTTP adapter errors and provider response bodies are never retained. +#[derive(Debug, Error, PartialEq, Eq)] +#[non_exhaustive] +pub enum EmaError { + #[error("invalid EMA exchange request: {0}")] + InvalidRequest(&'static str), + #[error("invalid EMA response from {stage:?}: {message}")] + InvalidResponse { + stage: EmaExchangeStage, + message: &'static str, + }, + #[error("EMA request to {0:?} failed")] + RequestFailed(EmaExchangeStage), + #[error("{0:?}: invalid_grant")] + InvalidGrant(EmaExchangeStage), + #[error("{0:?}: insufficient_user_authentication")] + InsufficientUserAuthentication(EmaExchangeStage), + #[error("{stage:?} returned HTTP {status}: {code}")] + OAuthRejected { + stage: EmaExchangeStage, + status: u16, + code: &'static str, + }, +} + +async fn post_form( + http: &dyn OAuthHttpClient, + server: &EmaAuthorizationServer, + params: &[(&str, &str)], + stage: EmaExchangeStage, + grant_expires_at: Option, + now: impl Fn() -> Result + Sync, +) -> Result { + let response = tokio::time::timeout(DEFAULT_HTTP_TIMEOUT, async { + // Generate assertions at the request boundary, not when server configuration is built. + let client_assertion = match &server.client_authentication { + EmaClientAuthentication::JwtAssertion(provider) => { + let assertion = provider + .create_assertion(server) + .await + .map_err(|_| EmaError::RequestFailed(stage))?; + if assertion.0.trim().is_empty() { + return Err(EmaError::InvalidRequest( + "client assertion must not be empty", + )); + } + Some(assertion) + } + _ => None, + }; + let request = { + let mut form = url::form_urlencoded::Serializer::new(String::new()); + form.extend_pairs(params.iter().copied()); + let mut request = oauth2::http::Request::builder() + .method("POST") + .uri(&server.token_endpoint) + .header("content-type", "application/x-www-form-urlencoded") + .header("accept", "application/json"); + match &server.client_authentication { + EmaClientAuthentication::ClientSecretBasic(secret) => { + let client_id: String = + url::form_urlencoded::byte_serialize(server.client_id.as_bytes()).collect(); + let secret: String = + url::form_urlencoded::byte_serialize(secret.secret().as_bytes()).collect(); + let encoded = STANDARD.encode(format!("{client_id}:{secret}")); + let mut header = + oauth2::http::HeaderValue::from_str(&format!("Basic {encoded}")).map_err( + |_| EmaError::InvalidRequest("invalid client authentication header"), + )?; + header.set_sensitive(true); + request = request.header(oauth2::http::header::AUTHORIZATION, header); + } + EmaClientAuthentication::ClientSecretPost(secret) => { + form.append_pair("client_id", &server.client_id) + .append_pair("client_secret", secret.secret()); + } + EmaClientAuthentication::None | EmaClientAuthentication::JwtAssertion(_) => { + form.append_pair("client_id", &server.client_id); + } + } + if let Some(assertion) = client_assertion { + form.append_pair( + "client_assertion_type", + "urn:ietf:params:oauth:client-assertion-type:jwt-bearer", + ) + .append_pair("client_assertion", &assertion.0); + } + request + .body(form.finish().into_bytes()) + .map_err(|_| EmaError::InvalidRequest("invalid token endpoint URI"))? + }; + // An external signer may outlive the grant even when it meets the request deadline. + if let Some(expires_at) = grant_expires_at + && expires_at <= now()? + { + return Err(EmaError::InvalidRequest("ID-JAG expired before redemption")); + } + http.execute(OAuthHttpRequest::new( + request, + OAuthHttpRedirectPolicy::Stop, + )) + .await + .map_err(|_| EmaError::RequestFailed(stage)) + }) + .await + .map_err(|_| EmaError::RequestFailed(stage))??; + let invalid = |message| EmaError::InvalidResponse { stage, message }; + if response.body().len() > MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES { + return Err(invalid("response body too large")); + } + if !response.status().is_success() { + #[derive(Deserialize)] + struct OAuthError { + error: Option, + } + let error = serde_json::from_slice::(response.body()).ok(); + let code = match error.as_ref().and_then(|e| e.error.as_deref()) { + Some("invalid_grant") => return Err(EmaError::InvalidGrant(stage)), + Some("insufficient_user_authentication") => { + return Err(EmaError::InsufficientUserAuthentication(stage)); + } + Some("invalid_request") => "invalid_request", + Some("invalid_client") => "invalid_client", + Some("invalid_scope") => "invalid_scope", + Some("invalid_target") => "invalid_target", + Some("unauthorized_client") => "unauthorized_client", + Some("unsupported_grant_type") => "unsupported_grant_type", + Some("access_denied") => "access_denied", + Some("temporarily_unavailable") => "temporarily_unavailable", + Some("server_error") => "server_error", + _ => "OAuth token request rejected", + }; + return Err(EmaError::OAuthRejected { + stage, + status: response.status().as_u16(), + code, + }); + } + serde_json::from_slice(response.body()).map_err(|_| invalid("malformed token response")) +} + +fn validate_endpoint(value: &str) -> Result<(), EmaError> { + let url = Url::parse(value).map_err(|_| EmaError::InvalidRequest("invalid endpoint URL"))?; + let loopback = match url.host() { + Some(Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + Some(Host::Ipv4(ip)) => ip.is_loopback(), + Some(Host::Ipv6(ip)) => ip.is_loopback(), + None => false, + }; + if (url.scheme() != "https" && !(url.scheme() == "http" && loopback)) + || !url.username().is_empty() + || url.password().is_some() + || url.fragment().is_some() + { + return Err(EmaError::InvalidRequest( + "endpoint must use HTTPS or HTTP loopback without userinfo or fragments", + )); + } + Ok(()) +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum Resource { + Single(String), + Multiple(Vec), +} + +impl Resource { + fn is_exact(&self, expected: &str) -> bool { + match self { + Self::Single(value) => value == expected, + Self::Multiple(values) => values.as_slice() == [expected], + } + } +} + +#[derive(Deserialize)] +struct JwtHeader { + alg: String, + typ: Option, +} + +#[derive(Deserialize)] +struct IdJagClaims { + iss: String, + sub: String, + aud: Resource, + client_id: String, + jti: String, + exp: u64, + iat: u64, + resource: Resource, + scope: Option, + authorization_details: Option>, +} + +#[derive(Deserialize)] +struct IdJagResponse { + access_token: String, + issued_token_type: String, + token_type: String, + resource: Option, + scope: Option, + refresh_token: Option, + authorization_details: Option>, +} + +impl IdJagResponse { + fn validate( + &self, + request: &EmaExchangeRequest<'_>, + requested: &HashSet<&str>, + ) -> Result<(HashSet, u64), EmaError> { + let invalid = |message| EmaError::InvalidResponse { + stage: EmaExchangeStage::IdentityProvider, + message, + }; + if self.issued_token_type != ID_JAG_TOKEN_TYPE + || self.token_type != "N_A" + || self.refresh_token.is_some() + { + return Err(invalid("unsupported ID-JAG token type or refresh token")); + } + let mut parts = self.access_token.split('.'); + let (Some(header), Some(payload), Some(signature), None) = + (parts.next(), parts.next(), parts.next(), parts.next()) + else { + return Err(invalid("ID-JAG must be a compact signed JWT")); + }; + if header.is_empty() || payload.is_empty() || signature.is_empty() { + return Err(invalid("ID-JAG contains an empty JWT segment")); + } + let decode = |value| { + URL_SAFE_NO_PAD + .decode(value) + .map_err(|_| invalid("malformed ID-JAG encoding")) + }; + decode(signature)?; + let header: JwtHeader = serde_json::from_slice(&decode(header)?) + .map_err(|_| invalid("malformed ID-JAG header"))?; + let claims: IdJagClaims = serde_json::from_slice(&decode(payload)?) + .map_err(|_| invalid("malformed ID-JAG claims"))?; + if [&self.authorization_details, &claims.authorization_details] + .into_iter() + .any(|details| details.as_ref().is_some_and(|details| !details.is_empty())) + { + return Err(invalid("authorization_details is not supported")); + } + if header.alg.trim().is_empty() + || header.alg.eq_ignore_ascii_case("none") + || header.typ.as_deref() != Some("oauth-id-jag+jwt") + || claims.iss != request.idp.issuer + || !claims.aud.is_exact(&request.resource_as.issuer) + || claims.client_id != request.resource_as.client_id + || claims.sub.trim().is_empty() + || claims.jti.trim().is_empty() + { + return Err(invalid( + "ID-JAG type, issuer, audience, client, subject, or JWT ID mismatch", + )); + } + let now = unix_time()?; + if claims.exp <= now || claims.iat > now.saturating_add(60) { + return Err(invalid("expired or future-issued ID-JAG")); + } + if !claims.resource.is_exact(request.resource) + || self + .resource + .as_ref() + .is_some_and(|r| !r.is_exact(request.resource)) + { + return Err(invalid("ID-JAG resource mismatch")); + } + let parse = + |scope| parse_scope(scope).ok_or_else(|| invalid("malformed or duplicate scopes")); + let granted = match claims.scope.as_deref() { + Some(scope) => parse(scope)?, + None if requested.is_empty() => HashSet::new(), + None => return Err(invalid("ID-JAG is missing requested scope authorization")), + }; + if !requested.is_empty() && !granted.is_subset(requested) { + return Err(invalid("ID-JAG scope exceeds the request")); + } + match self.scope.as_deref() { + Some(scope) if parse(scope)? != granted => { + return Err(invalid("response scope differs from ID-JAG scope")); + } + None if !requested.is_empty() && granted != *requested => { + return Err(invalid("response omitted narrowed scope")); + } + _ => {} + } + Ok((granted.into_iter().map(str::to_owned).collect(), claims.exp)) + } +} + +fn parse_scope(scope: &str) -> Option> { + let scopes: HashSet<_> = scope.split(' ').collect(); + (scopes.iter().all(|s| is_scope_token(s)) && scopes.len() == scope.split(' ').count()) + .then_some(scopes) +} + +fn is_scope_token(scope: &str) -> bool { + !scope.is_empty() + && scope + .bytes() + .all(|b| matches!(b, b'!' | b'#'..=b'[' | b']'..=b'~')) +} + +fn unix_time() -> Result { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_secs()) + .map_err(|_| EmaError::InvalidRequest("system clock precedes the Unix epoch")) +} + +#[derive(Deserialize)] +struct ResourceTokenResponse { + access_token: String, + token_type: String, + expires_in: Option, + resource: Option, + scope: Option, + refresh_token: Option, + authorization_details: Option>, +} + +impl ResourceTokenResponse { + fn validate( + self, + resource: &str, + granted: HashSet, + ) -> Result { + let invalid = |message| EmaError::InvalidResponse { + stage: EmaExchangeStage::ResourceAuthorizationServer, + message, + }; + if self + .authorization_details + .as_ref() + .is_some_and(|details| !details.is_empty()) + { + return Err(invalid("authorization_details is not supported")); + } + if !self.token_type.eq_ignore_ascii_case("bearer") + || self.access_token.trim().is_empty() + || self.refresh_token.is_some() + || self.expires_in == Some(0) + { + return Err(invalid( + "invalid bearer token, lifetime, or unexpected refresh token", + )); + } + // The resource need not be echoed, but must agree with the ID-JAG if present. + if self + .resource + .as_ref() + .is_some_and(|r| !r.is_exact(resource)) + { + return Err(invalid("access token resource mismatch")); + } + // An omitted scope retains the authority carried by the assertion. + let scopes = if let Some(scope) = self.scope.as_deref() { + let scopes = + parse_scope(scope).ok_or_else(|| invalid("malformed or duplicate scopes"))?; + if !scopes.iter().all(|s| granted.contains(*s)) { + return Err(invalid("access token scope exceeds ID-JAG scope")); + } + scopes.into_iter().map(str::to_owned).collect() + } else { + granted + }; + Ok(EmaAccessToken { + access_token: AccessToken::new(self.access_token), + expires_in: self.expires_in.map(Duration::from_secs), + scopes, + }) + } +} + +#[cfg(test)] +#[path = "enterprise_tests.rs"] +mod tests; diff --git a/crates/rmcp/src/transport/auth/enterprise_tests.rs b/crates/rmcp/src/transport/auth/enterprise_tests.rs new file mode 100644 index 000000000..31a63c1b2 --- /dev/null +++ b/crates/rmcp/src/transport/auth/enterprise_tests.rs @@ -0,0 +1,1145 @@ +use std::{ + collections::BTreeMap, + sync::{Arc, Mutex}, + time::{Duration, SystemTime}, +}; + +use base64::{ + Engine, + engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}, +}; +use oauth2::{ClientSecret, HttpResponse, RefreshToken}; +use serde_json::{Value, json}; + +use super::{ + EmaExchangeStage::{IdentityProvider as Idp, ResourceAuthorizationServer as ResourceServer}, + *, +}; +use crate::transport::auth::{OAuthHttpClientError, OAuthHttpClientFuture, OAuthHttpRequest}; + +const IDP: &str = "https://idp.example?private-query"; +const AS: &str = "https://as.example?private-query"; +const RESOURCE: &str = "https://mcp.example?private-query"; +const IDP_TOKEN: &str = "https://idp.example/token?private-query"; +const AS_TOKEN: &str = "https://as.example/token?private-query"; +const BAD_SCOPES: &[&str] = &["", "files\tread", "\"", "\\", "\0", "读"]; + +#[derive(Default)] +struct MockHttp { + requests: Mutex>, + response: Mutex>>, +} + +impl MockHttp { + fn new(status: u16, body: Value) -> Self { + Self { + response: Mutex::new(Some(Ok(oauth2::http::Response::builder() + .status(status) + .body(serde_json::to_vec(&body).unwrap()) + .unwrap()))), + ..Self::default() + } + } +} + +impl OAuthHttpClient for MockHttp { + fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + self.requests.lock().unwrap().push(request); + let response = self.response.lock().unwrap().take(); + let response = response.expect("unexpected HTTP"); + Box::pin(async move { response }) + } +} + +fn jwt(header: Value, claims: &Value) -> String { + format!( + "{}.{}.{}", + URL_SAFE_NO_PAD.encode(header.to_string()), + URL_SAFE_NO_PAD.encode(claims.to_string()), + URL_SAFE_NO_PAD.encode(b"synthetic-signature") + ) +} + +fn claims() -> Value { + let now = SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs(); + json!({"iss":IDP,"aud":AS,"sub":"user","client_id":"mcp", + "jti":"jag-id","iat":now,"exp":now + 3600,"resource":RESOURCE,"scope":"files.read"}) +} + +fn jag(claims: &Value) -> Value { + let mut body = json!({"access_token":jwt(json!({"alg":"ES256","typ":"oauth-id-jag+jwt"}), claims), + "issued_token_type":ID_JAG_TOKEN_TYPE,"token_type":"N_A","resource":claims["resource"]}); + if let Some(scope) = claims.get("scope") { + body["scope"] = scope.clone(); + } + body +} + +async fn exchange(idp: &MockHttp, scopes: &str) -> Result { + let scopes = scopes + .split_ascii_whitespace() + .map(str::to_owned) + .collect::>(); + let refresh = RefreshToken::new("refresh-token".into()); + let resource_as = EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"); + assert!(!format!("{resource_as:?}").contains("private-query")); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + resource_as, + RESOURCE, + &refresh, + ) + .with_scopes(scopes); + assert!(!format!("{request:?}").contains(refresh.secret())); + assert!(!format!("{request:?}").contains("private-query")); + let future = request.exchange_id_jag(idp); + fn is_send(_: &T) {} + is_send(&future); + future.await +} + +fn form(request: &OAuthHttpRequest) -> BTreeMap { + assert!(!request.request.headers().contains_key("authorization")); + authenticated_form(request) +} + +fn authenticated_form(request: &OAuthHttpRequest) -> BTreeMap { + assert_eq!(request.request.method(), "POST"); + let headers = request.request.headers(); + assert_eq!(headers["content-type"], "application/x-www-form-urlencoded"); + assert_eq!(headers["accept"], "application/json"); + assert_eq!(request.redirect_policy, OAuthHttpRedirectPolicy::Stop); + assert_eq!(request.timeout, Some(Duration::from_secs(30))); + let pairs = url::form_urlencoded::parse(request.request.body()) + .into_owned() + .collect::>(); + let fields = pairs.iter().cloned().collect::>(); + assert_eq!(pairs.len(), fields.len(), "duplicate form fields"); + fields +} + +#[tokio::test] +async fn refresh_exchange_preserves_exact_forms_and_signed_narrowing() { + for scopes in ["files.read files.write", ""] { + let response = jag(&claims()); + let idp = MockHttp::new(200, response.clone()); + let result = exchange(&idp, scopes).await.unwrap(); + assert_eq!( + result.assertion().secret(), + response["access_token"].as_str().unwrap() + ); + assert_eq!(result.scopes(), &HashSet::from(["files.read".to_owned()])); + assert!(!format!("{result:?}").contains(result.assertion().secret())); + assert!(!format!("{result:?}").contains("private-query")); + let requests = idp.requests.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].request.uri(), IDP_TOKEN); + let mut fields = form(&requests[0]); + assert_eq!( + fields.remove("scope"), + (!scopes.is_empty()).then(|| scopes.to_owned()) + ); + assert_eq!( + serde_json::to_value(fields).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:token-exchange", + "requested_token_type":ID_JAG_TOKEN_TYPE,"subject_token":"refresh-token", + "subject_token_type":"urn:ietf:params:oauth:token-type:refresh_token", + "audience":AS,"resource":RESOURCE,"client_id":"idp" + }) + ); + } +} + +#[tokio::test] +async fn scopes_may_be_omitted_but_never_widened() { + for (requested, signed, echoed, valid) in [ + ("", Some("files.read"), None, true), + ("", None, None, true), + ("files.read", Some("files.read"), None, true), + ("files.read files.write", Some("files.read"), None, false), + ( + "files.read", + Some("files.read files.write"), + Some("files.read files.write"), + false, + ), + ("files.read", None, None, false), + ("", Some("files.read"), Some("files.write"), false), + ("", Some(" \t"), None, false), + ("", Some("files.read files.read"), None, false), + ("", Some("files.read"), Some(" \t"), false), + ("", Some("files.read"), Some("files.read files.read"), false), + ] { + let mut claims = claims(); + claims.as_object_mut().unwrap().remove("scope"); + if let Some(scope) = signed { + claims["scope"] = json!(scope); + } + let mut response = jag(&claims); + response.as_object_mut().unwrap().remove("scope"); + if let Some(scope) = echoed { + response["scope"] = json!(scope); + } + let result = exchange(&MockHttp::new(200, response), requested).await; + assert_eq!( + result.is_ok(), + valid, + "requested={requested:?}, signed={signed:?}, echoed={echoed:?}" + ); + if let Ok(result) = result { + let expected = signed.into_iter().map(str::to_owned).collect(); + assert_eq!(result.scopes(), &expected); + } + } +} + +#[tokio::test] +async fn unsupported_authorization_details_are_rejected_at_each_stage() { + const SECRET: &str = "authorization-details-secret"; + for (location, stage) in [ + ("claims", Idp), + ("idp response", Idp), + ("resource response", ResourceServer), + ] { + for (case, details, valid) in [ + ("absent", None, true), + ("null", Some(Value::Null), true), + ("empty", Some(json!([])), true), + ( + "additional authority", + Some( + json!([{"type":SECRET,"locations":["https://other.example"],"actions":["write"]}]), + ), + false, + ), + ("invalid member", Some(json!([null])), false), + ("object", Some(json!({"type":SECRET})), false), + ("string", Some(json!(SECRET)), false), + ] { + let mut claims = claims(); + let mut response = bearer(); + // Unrelated extension fields remain compatible at every boundary. + claims["vendor_extension"] = json!(SECRET); + response["vendor_extension"] = json!(SECRET); + if location == "claims" + && let Some(details) = &details + { + claims["authorization_details"] = details.clone(); + } + let mut idp_response = jag(&claims); + idp_response["vendor_extension"] = json!(SECRET); + if let Some(details) = details { + match location { + "idp response" => idp_response["authorization_details"] = details, + "resource response" => response["authorization_details"] = details, + _ => {} + } + } + let idp = MockHttp::new(200, idp_response); + let resource = MockHttp::new(200, response); + let result = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .with_scopes(["files.read"]) + .exchange(&idp, &resource) + .await; + assert_eq!(result.is_ok(), valid, "{location}: {case}"); + if let Err(error) = result { + assert!( + matches!(error, EmaError::InvalidResponse { stage: actual, .. } if actual == stage) + ); + assert!(!format!("{error:?} {error}").contains(SECRET)); + } + assert_eq!(idp.requests.lock().unwrap().len(), 1); + assert_eq!( + resource.requests.lock().unwrap().len(), + usize::from(valid || stage == ResourceServer), + "{location}: {case}" + ); + } + } +} + +#[tokio::test] +async fn invalid_jags_never_escape_validation() { + let original = claims(); + let mut cases = Vec::new(); + let invalid = json!({"iss":"https://other.example","aud":[AS,"other"], + "client_id":"other","sub":"","jti":" \t","exp":0,"iat":u64::MAX,"resource":[RESOURCE,"other"]}); + for (key, value) in invalid.as_object().unwrap() { + let mut changed = original.clone(); + changed[key] = value.clone(); + cases.push(jag(&changed)); + } + for scope in BAD_SCOPES { + let mut changed = original.clone(); + changed["scope"] = json!(scope); + cases.push(jag(&changed)); + } + for header in [ + json!({"alg":"ES256","typ":"JWT"}), + json!({"alg":"ES256"}), + json!({"alg":"none","typ":"oauth-id-jag+jwt"}), + ] { + let mut response = jag(&original); + response["access_token"] = json!(jwt(header, &original)); + cases.push(response); + } + let invalid = json!({"issued_token_type":"Bearer","token_type":"Bearer", + "refresh_token":"unsupported","resource":"https://other.example","access_token":"a.b.c.d"}); + for (key, value) in invalid.as_object().unwrap() { + let mut response = jag(&original); + response[key] = value.clone(); + cases.push(response); + } + for signature in ["", "signature", "not+base64url", "c2ln="] { + let mut response = jag(&original); + let assertion = response["access_token"].as_str().unwrap(); + let (signed, _) = assertion.rsplit_once('.').unwrap(); + response["access_token"] = json!(format!("{signed}.{signature}")); + cases.push(response); + } + for response in cases { + let result = exchange(&MockHttp::new(200, response), "files.read").await; + assert!(matches!( + result, + Err(EmaError::InvalidResponse { stage: Idp, .. }) + )); + } +} + +#[tokio::test] +async fn invalid_inputs_fail_before_http() { + // Each server occupies issuer, token endpoint, and client ID slots. + let original = [IDP, IDP_TOKEN, "idp", AS, AS_TOKEN, "mcp", RESOURCE, "rt"]; + let mut cases = Vec::new(); + for index in [0, 1, 3, 4, 6] { + for value in [ + "invalid", + "http://idp.example/token", + "https://user:pass@idp.example/token", + "https://idp.example/token#fragment", + ] { + let mut fields = original; + fields[index] = value; + cases.push((fields, vec![])); + } + } + for (index, value) in [(0, AS), (2, " "), (5, ""), (7, " \t")] { + let mut fields = original; + fields[index] = value; + cases.push((fields, vec![])); + } + for scope in BAD_SCOPES.iter().copied().chain(["files.read files.write"]) { + cases.push((original, vec![scope])); + } + cases.push((original, vec!["files.read", "files.read"])); + for (fields, scopes) in cases { + let http = MockHttp::default(); + let result = EmaExchangeRequest::new( + EmaAuthorizationServer::new(fields[0], fields[1], fields[2]), + EmaAuthorizationServer::new(fields[3], fields[4], fields[5]), + fields[6], + &RefreshToken::new(fields[7].into()), + ) + .with_scopes(scopes) + .exchange_id_jag(&http) + .await; + assert!(matches!(result, Err(EmaError::InvalidRequest(_)))); + assert!(http.requests.lock().unwrap().is_empty()); + } +} + +#[tokio::test] +async fn errors_and_redirects_cannot_reflect_credentials() { + const SECRET: &str = "secret-error-sentinel"; + for (status, code) in [ + (400, "invalid_grant"), + (400, "insufficient_user_authentication"), + (400, "invalid_client"), + (400, SECRET), + (302, SECRET), + ] { + let failure = json!({"error":code,"error_description":SECRET}); + let error = exchange(&MockHttp::new(status, failure), "") + .await + .unwrap_err(); + match code { + "invalid_grant" => assert_eq!(error, EmaError::InvalidGrant(Idp)), + "insufficient_user_authentication" => { + assert_eq!(error, EmaError::InsufficientUserAuthentication(Idp)) + } + _ => assert!( + matches!(error, EmaError::OAuthRejected {stage: Idp, status: actual, ..} if actual == status) + ), + } + assert!(!format!("{error:?} {error}").contains(SECRET)); + } + for adapter_failure in [false, true] { + let failure = MockHttp::new( + 200, + json!({"access_token":SECRET,"issued_token_type":SECRET}), + ); + if adapter_failure { + *failure.response.lock().unwrap() = Some(Err(SECRET.into())); + } + let error = exchange(&failure, "").await.unwrap_err(); + assert!(!format!("{error:?} {error}").contains(SECRET)); + assert!(std::error::Error::source(&error).is_none()); + if adapter_failure { + assert_eq!(error, EmaError::RequestFailed(Idp)); + } + } + let mut oversized = jag(&claims()); + oversized["ignored"] = json!("x".repeat(1024 * 1024)); + let error = exchange(&MockHttp::new(200, oversized), "") + .await + .unwrap_err(); + assert!(matches!( + error, + EmaError::InvalidResponse { stage: Idp, .. } + )); +} + +fn bearer() -> Value { + json!({"access_token":"resource-token","token_type":"Bearer","expires_in":300}) +} + +async fn redeem(http: &MockHttp) -> Result { + exchange( + &MockHttp::new(200, jag(&claims())), + "files.read files.write", + ) + .await? + .exchange(http) + .await +} + +#[tokio::test] +async fn full_exchange_uses_separate_clients_and_only_the_narrowed_assertion() { + for valid in [true, false] { + let mut claims = claims(); + if !valid { + claims["client_id"] = json!("other"); + } + let response = jag(&claims); + let idp = MockHttp::new(200, response.clone()); + let resource_as = MockHttp::new(200, bearer()); + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"), + RESOURCE, + &refresh, + ) + .with_scopes(["files.read", "files.write"]); + let future = request.exchange(&idp, &resource_as); + fn is_send(_: &T) {} + is_send(&future); + let result = future.await; + assert_eq!(idp.requests.lock().unwrap().len(), 1); + let requests = resource_as.requests.lock().unwrap(); + assert_eq!(requests.len(), usize::from(valid)); + if !valid { + assert!(matches!( + result, + Err(EmaError::InvalidResponse { stage: Idp, .. }) + )); + continue; + } + let token = result.unwrap(); + assert_eq!(token.access_token.secret(), "resource-token"); + assert_eq!(token.expires_in, Some(Duration::from_secs(300))); + assert_eq!(token.scopes, HashSet::from(["files.read".to_owned()])); + assert!(!format!("{token:?}").contains(token.access_token.secret())); + assert_eq!(requests[0].request.uri(), AS_TOKEN); + assert_eq!( + serde_json::to_value(form(&requests[0])).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:jwt-bearer", + "assertion":response["access_token"],"client_id":"mcp" + }) + ); + } +} + +#[tokio::test] +async fn expired_grants_are_not_redeemed() { + let mut grant = exchange(&MockHttp::new(200, jag(&claims())), "") + .await + .unwrap(); + grant.expires_at = 0; + let resource_as = MockHttp::default(); + assert!(matches!( + grant.exchange(&resource_as).await, + Err(EmaError::InvalidRequest(_)) + )); + assert!(resource_as.requests.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn bearer_responses_cannot_change_resource_or_widen_scope() { + let mut cases = vec![ + ("scope", json!("files.read"), true), + ("scope", json!("files.admin"), false), + ("scope", json!("files.read files.write"), false), + ("scope", json!("files.read files.read"), false), + ("resource", json!(RESOURCE), true), + ("resource", json!([RESOURCE]), true), + ("resource", json!([RESOURCE, "other"]), false), + ("resource", json!("https://mcp.example"), false), + ("expires_in", json!(0), false), + ("expires_in", Value::Null, true), + ("refresh_token", json!("unsupported"), false), + ("token_type", json!("N_A"), false), + ("token_type", json!("bearer"), true), + ("access_token", json!(" \t"), false), + ]; + cases.extend( + BAD_SCOPES + .iter() + .map(|scope| ("scope", json!(scope), false)), + ); + for (key, value, valid) in cases { + let mut token = bearer(); + if value.is_null() { + token.as_object_mut().unwrap().remove(key); + } else { + token[key] = value; + } + let result = redeem(&MockHttp::new(200, token)).await; + assert_eq!(result.is_ok(), valid, "{key}"); + if key == "expires_in" && valid { + assert_eq!(result.unwrap().expires_in, None); + } else if !valid { + assert!(matches!( + result, + Err(EmaError::InvalidResponse { + stage: ResourceServer, + .. + }) + )); + } + } +} + +#[tokio::test] +async fn resource_scope_narrowing_is_reported_without_logging_scope_values() { + for read in ["files.read", "private-scope-sentinel"] { + let requested = format!("{read} files.write"); + let mut claims = claims(); + claims["scope"] = json!(requested); + let grant = exchange(&MockHttp::new(200, jag(&claims)), &requested) + .await + .unwrap(); + let mut response = bearer(); + response["scope"] = json!(read); + let token = grant.exchange(&MockHttp::new(200, response)).await.unwrap(); + assert_eq!(token.scopes, HashSet::from([read.to_owned()])); + assert!(!format!("{token:?}").contains(read)); + } +} + +#[tokio::test] +async fn bearer_scope_may_be_omitted_but_not_added_to_an_unscoped_grant() { + for scope in [None, Some("files.read")] { + let mut claims = claims(); + claims.as_object_mut().unwrap().remove("scope"); + let grant = exchange(&MockHttp::new(200, jag(&claims)), "") + .await + .unwrap(); + let mut token = bearer(); + if let Some(scope) = scope { + token["scope"] = json!(scope); + } + let result = grant.exchange(&MockHttp::new(200, token)).await; + assert_eq!(result.is_ok(), scope.is_none()); + if let Ok(token) = result { + assert!(token.scopes.is_empty()); + } + } +} + +#[tokio::test] +async fn resource_errors_are_staged_and_never_reflect_credentials() { + const SECRET: &str = "secret-resource-error-sentinel"; + for (status, code) in [ + (400, "invalid_grant"), + (400, "insufficient_user_authentication"), + (400, "invalid_client"), + (302, SECRET), + (500, SECRET), + ] { + let http = MockHttp::new(status, json!({"error":code,"error_description":SECRET})); + let error = redeem(&http).await.unwrap_err(); + match code { + "invalid_grant" => assert_eq!(error, EmaError::InvalidGrant(ResourceServer)), + "insufficient_user_authentication" => assert_eq!( + error, + EmaError::InsufficientUserAuthentication(ResourceServer) + ), + _ => assert!( + matches!(error, EmaError::OAuthRejected {stage: ResourceServer, status: actual, ..} if actual == status) + ), + } + assert!(!format!("{error:?} {error}").contains(SECRET)); + assert!(std::error::Error::source(&error).is_none()); + assert_eq!(http.requests.lock().unwrap().len(), 1); + } + let mut oversized = bearer(); + oversized["ignored"] = json!("x".repeat(1024 * 1024)); + for (body, adapter_failure) in [ + (json!({"access_token":SECRET,"expires_in":SECRET}), false), + (oversized, false), + (Value::Null, true), + ] { + let http = MockHttp::new(200, body); + if adapter_failure { + *http.response.lock().unwrap() = Some(Err(SECRET.into())); + } + let error = redeem(&http).await.unwrap_err(); + assert!(!format!("{error:?} {error}").contains(SECRET)); + assert!(std::error::Error::source(&error).is_none()); + if adapter_failure { + assert_eq!(error, EmaError::RequestFailed(ResourceServer)); + } else { + assert!(matches!( + error, + EmaError::InvalidResponse { + stage: ResourceServer, + .. + } + )); + } + } +} + +#[tokio::test(start_paused = true)] +async fn both_exchanges_enforce_the_timeout_when_the_adapter_does_not() { + struct PendingHttp; + impl OAuthHttpClient for PendingHttp { + fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + Box::pin(std::future::pending()) + } + } + for stage in [Idp, ResourceServer] { + let ready = MockHttp::new(200, jag(&claims())); + let (idp, resource): (&dyn OAuthHttpClient, &dyn OAuthHttpClient) = match stage { + Idp => (&PendingHttp, &ready), + _ => (&ready, &PendingHttp), + }; + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"), + RESOURCE, + &refresh, + ); + let start = tokio::time::Instant::now(); + let result = tokio::time::timeout(Duration::from_secs(31), request.exchange(idp, resource)) + .await + .expect("the SDK must enforce its own deadline"); + assert_eq!(result.unwrap_err(), EmaError::RequestFailed(stage)); + assert_eq!(start.elapsed(), Duration::from_secs(30)); + assert_eq!( + ready.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + } +} + +#[derive(Clone, Copy)] +enum AssertionBehavior { + Success, + Error, + Empty(&'static str), + Delay(Duration), + Pending, +} + +struct RecordingAssertionProvider { + calls: Mutex>, + behavior: AssertionBehavior, +} + +impl RecordingAssertionProvider { + fn new(behavior: AssertionBehavior) -> Arc { + Arc::new(Self { + calls: Mutex::new(Vec::new()), + behavior, + }) + } +} + +fn client_assertion(client_id: &str, issuer: &str, sequence: usize) -> String { + jwt( + json!({"alg":"ES256","typ":"client-authentication+jwt"}), + &json!({"iss":client_id,"sub":client_id,"aud":issuer, + "exp":4102444800_u64,"jti":format!("client-assertion-{sequence}")}), + ) +} + +#[async_trait::async_trait] +impl EmaClientAssertionProvider for RecordingAssertionProvider { + async fn create_assertion( + &self, + server: &EmaAuthorizationServer, + ) -> Result> { + let sequence = { + let mut calls = self.calls.lock().unwrap(); + calls.push(( + server.issuer.clone(), + server.token_endpoint.clone(), + server.client_id.clone(), + )); + calls.len() + }; + match self.behavior { + AssertionBehavior::Error => return Err("client-assertion-provider-secret".into()), + AssertionBehavior::Empty(value) => return Ok(EmaClientAssertion::new(value.into())), + AssertionBehavior::Delay(delay) => tokio::time::sleep(delay).await, + AssertionBehavior::Pending => std::future::pending::<()>().await, + AssertionBehavior::Success => {} + } + Ok(EmaClientAssertion::new(client_assertion( + &server.client_id, + &server.issuer, + sequence, + ))) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum AuthenticationMethod { + Public, + Basic, + Post, + Jwt, +} + +impl AuthenticationMethod { + fn configure( + self, + secret: &str, + provider: &Arc, + ) -> EmaClientAuthentication { + match self { + Self::Public => EmaClientAuthentication::None, + Self::Basic => { + EmaClientAuthentication::ClientSecretBasic(ClientSecret::new(secret.into())) + } + Self::Post => { + EmaClientAuthentication::ClientSecretPost(ClientSecret::new(secret.into())) + } + Self::Jwt => EmaClientAuthentication::JwtAssertion(provider.clone()), + } + } +} + +fn remove_client_authentication( + request: &OAuthHttpRequest, + method: AuthenticationMethod, + client_id: &str, + secret: &str, + encoded_basic: &str, + assertion: &str, +) -> BTreeMap { + let mut fields = authenticated_form(request); + let authorization = request.request.headers().get("authorization"); + if method == AuthenticationMethod::Basic { + let authorization = authorization.expect("Basic authentication must use a header"); + assert!(authorization.is_sensitive()); + let encoded = authorization + .to_str() + .unwrap() + .strip_prefix("Basic ") + .unwrap(); + assert_eq!(STANDARD.decode(encoded).unwrap(), encoded_basic.as_bytes()); + assert!(!fields.contains_key("client_id")); + } else { + assert!(authorization.is_none()); + assert_eq!(fields.remove("client_id").as_deref(), Some(client_id)); + } + assert_eq!( + fields.remove("client_secret").as_deref(), + (method == AuthenticationMethod::Post).then_some(secret) + ); + assert_eq!( + fields.remove("client_assertion_type").as_deref(), + (method == AuthenticationMethod::Jwt) + .then_some("urn:ietf:params:oauth:client-assertion-type:jwt-bearer") + ); + assert_eq!( + fields.remove("client_assertion").as_deref(), + (method == AuthenticationMethod::Jwt).then_some(assertion) + ); + fields +} + +#[tokio::test] +async fn client_authentication_is_endpoint_specific_and_never_mixed_with_grants() { + const METHODS: [AuthenticationMethod; 4] = [ + AuthenticationMethod::Public, + AuthenticationMethod::Basic, + AuthenticationMethod::Post, + AuthenticationMethod::Jwt, + ]; + // Colons, plus signs, spaces, percent signs, and non-ASCII bytes must be + // form-encoded individually before the HTTP Basic username/password join. + const IDP_CLIENT: &str = "idp:+ %é"; + const AS_CLIENT: &str = "mcp:+ %é"; + const IDP_SECRET: &str = "idp-secret:+ %é"; + const AS_SECRET: &str = "as-secret:+ %é"; + for idp_method in METHODS { + for resource_method in METHODS { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let mut claims = claims(); + claims["client_id"] = json!(AS_CLIENT); + let response = jag(&claims); + let idp = MockHttp::new(200, response.clone()); + let resource = MockHttp::new(200, bearer()); + let refresh = RefreshToken::new("refresh-token".into()); + let token = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, IDP_CLIENT) + .with_client_authentication(idp_method.configure(IDP_SECRET, &provider)), + EmaAuthorizationServer::new(AS, AS_TOKEN, AS_CLIENT) + .with_client_authentication(resource_method.configure(AS_SECRET, &provider)), + RESOURCE, + &refresh, + ) + .exchange(&idp, &resource) + .await + .unwrap(); + assert_eq!(token.access_token.secret(), "resource-token"); + let idp_requests = idp.requests.lock().unwrap(); + let resource_requests = resource.requests.lock().unwrap(); + assert_eq!(idp_requests.len(), 1); + assert_eq!(resource_requests.len(), 1); + assert_eq!(idp_requests[0].request.uri(), IDP_TOKEN); + assert_eq!(resource_requests[0].request.uri(), AS_TOKEN); + let idp_fields = remove_client_authentication( + &idp_requests[0], + idp_method, + IDP_CLIENT, + IDP_SECRET, + "idp%3A%2B+%25%C3%A9:idp-secret%3A%2B+%25%C3%A9", + &client_assertion(IDP_CLIENT, IDP, 1), + ); + assert_eq!( + serde_json::to_value(idp_fields).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:token-exchange", + "requested_token_type":ID_JAG_TOKEN_TYPE,"subject_token":"refresh-token", + "subject_token_type":"urn:ietf:params:oauth:token-type:refresh_token", + "audience":AS,"resource":RESOURCE + }) + ); + let resource_fields = remove_client_authentication( + &resource_requests[0], + resource_method, + AS_CLIENT, + AS_SECRET, + "mcp%3A%2B+%25%C3%A9:as-secret%3A%2B+%25%C3%A9", + &client_assertion( + AS_CLIENT, + AS, + 1 + usize::from(idp_method == AuthenticationMethod::Jwt), + ), + ); + assert_eq!( + serde_json::to_value(resource_fields).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:jwt-bearer", + "assertion":response["access_token"] + }) + ); + let mut expected_calls = Vec::new(); + if idp_method == AuthenticationMethod::Jwt { + expected_calls.push((IDP.into(), IDP_TOKEN.into(), IDP_CLIENT.into())); + } + if resource_method == AuthenticationMethod::Jwt { + expected_calls.push((AS.into(), AS_TOKEN.into(), AS_CLIENT.into())); + } + assert_eq!(*provider.calls.lock().unwrap(), expected_calls); + } + } +} + +#[tokio::test(start_paused = true)] +async fn client_assertions_are_fresh_for_every_request_and_delayed_redemption() { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let idp_server = EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(EmaClientAuthentication::JwtAssertion(provider.clone())); + let resource_server = EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(EmaClientAuthentication::JwtAssertion(provider.clone())); + let refresh = RefreshToken::new("refresh-token".into()); + let mut assertions = HashSet::new(); + for attempt in 0..2 { + let idp = MockHttp::new(200, jag(&claims())); + let resource = MockHttp::new(200, bearer()); + let grant = EmaExchangeRequest::new( + idp_server.clone(), + resource_server.clone(), + RESOURCE, + &refresh, + ) + .exchange_id_jag(&idp) + .await + .unwrap(); + assert_eq!(provider.calls.lock().unwrap().len(), attempt * 2 + 1); + tokio::time::sleep(Duration::from_secs(60)).await; + assert_eq!(provider.calls.lock().unwrap().len(), attempt * 2 + 1); + grant.exchange(&resource).await.unwrap(); + assert_eq!(provider.calls.lock().unwrap().len(), attempt * 2 + 2); + for http in [&idp, &resource] { + let requests = http.requests.lock().unwrap(); + let mut fields = form(&requests[0]); + assert!(assertions.insert(fields.remove("client_assertion").unwrap())); + } + } + assert_eq!(assertions.len(), 4); +} + +#[tokio::test] +async fn invalid_static_client_secrets_fail_before_any_http_or_signing() { + for stage in [Idp, ResourceServer] { + for method in [AuthenticationMethod::Basic, AuthenticationMethod::Post] { + for secret in ["", " \t"] { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let valid = EmaClientAuthentication::JwtAssertion(provider.clone()); + let invalid = method.configure(secret, &provider); + let (idp_auth, resource_auth) = match stage { + Idp => (invalid, valid), + _ => (valid, invalid), + }; + let idp = MockHttp::default(); + let resource = MockHttp::default(); + let error = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(idp_auth), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(resource_auth), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange(&idp, &resource) + .await + .unwrap_err(); + assert!(matches!(error, EmaError::InvalidRequest(_))); + assert!(idp.requests.lock().unwrap().is_empty()); + assert!(resource.requests.lock().unwrap().is_empty()); + assert!(provider.calls.lock().unwrap().is_empty()); + } + } + } +} + +#[tokio::test] +async fn a_grant_that_expires_while_signing_is_never_sent_to_the_resource_server() { + use std::sync::atomic::{AtomicU64, Ordering}; + + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let idp = MockHttp::new(200, jag(&claims())); + let resource = MockHttp::default(); + let mut grant = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(EmaClientAuthentication::JwtAssertion(provider.clone())), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange_id_jag(&idp) + .await + .unwrap(); + grant.expires_at = 100; + // The clock crosses expiration between the checks before and after signing. + let now = AtomicU64::new(99); + let error = grant + .exchange_with_clock(&resource, || Ok(now.fetch_add(1, Ordering::Relaxed))) + .await + .unwrap_err(); + assert!(matches!(error, EmaError::InvalidRequest(_))); + assert_eq!(provider.calls.lock().unwrap().len(), 1); + assert!(resource.requests.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn assertion_provider_failures_and_empty_results_never_reach_http_or_escape_errors() { + for stage in [Idp, ResourceServer] { + for behavior in [ + AssertionBehavior::Error, + AssertionBehavior::Empty(""), + AssertionBehavior::Empty(" \t"), + ] { + let provider = RecordingAssertionProvider::new(behavior); + let auth = EmaClientAuthentication::JwtAssertion(provider.clone()); + let (idp_auth, resource_auth) = match stage { + Idp => (auth, EmaClientAuthentication::None), + _ => (EmaClientAuthentication::None, auth), + }; + let idp = MockHttp::new(200, jag(&claims())); + let resource = MockHttp::default(); + let error = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(idp_auth), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(resource_auth), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange(&idp, &resource) + .await + .unwrap_err(); + if matches!(behavior, AssertionBehavior::Error) { + assert_eq!(error, EmaError::RequestFailed(stage)); + } else { + assert!(matches!(error, EmaError::InvalidRequest(_))); + } + assert!(!format!("{error:?} {error}").contains("client-assertion-provider-secret")); + assert!(std::error::Error::source(&error).is_none()); + assert_eq!(provider.calls.lock().unwrap().len(), 1); + assert_eq!( + idp.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + assert!(resource.requests.lock().unwrap().is_empty()); + } + } +} + +#[tokio::test(start_paused = true)] +async fn signing_and_http_share_one_deadline_at_both_endpoints() { + struct DelayedHttp(Mutex>); + impl OAuthHttpClient for DelayedHttp { + fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + self.0.lock().unwrap().push(request); + Box::pin(async { + tokio::time::sleep(Duration::from_secs(10)).await; + Err("http-adapter-secret".into()) + }) + } + } + for stage in [Idp, ResourceServer] { + for behavior in [ + AssertionBehavior::Pending, + AssertionBehavior::Delay(Duration::from_secs(25)), + ] { + let provider = RecordingAssertionProvider::new(behavior); + let auth = EmaClientAuthentication::JwtAssertion(provider.clone()); + let (idp_auth, resource_auth) = match stage { + Idp => (auth, EmaClientAuthentication::None), + _ => (EmaClientAuthentication::None, auth), + }; + let ready = MockHttp::new(200, jag(&claims())); + let delayed = DelayedHttp(Mutex::new(Vec::new())); + let (idp, resource): (&dyn OAuthHttpClient, &dyn OAuthHttpClient) = match stage { + Idp => (&delayed, &ready), + _ => (&ready, &delayed), + }; + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(idp_auth), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(resource_auth), + RESOURCE, + &refresh, + ); + let start = tokio::time::Instant::now(); + let result = + tokio::time::timeout(Duration::from_secs(31), request.exchange(idp, resource)) + .await + .expect("signing must share the SDK's token request deadline"); + assert_eq!(result.unwrap_err(), EmaError::RequestFailed(stage)); + assert_eq!(start.elapsed(), Duration::from_secs(30)); + assert_eq!(provider.calls.lock().unwrap().len(), 1); + assert_eq!( + delayed.0.lock().unwrap().len(), + usize::from(matches!(behavior, AssertionBehavior::Delay(_))) + ); + assert_eq!( + ready.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + } + } +} + +#[tokio::test] +async fn rejected_client_authentication_never_falls_back_or_retries() { + for stage in [Idp, ResourceServer] { + for method in [ + AuthenticationMethod::Basic, + AuthenticationMethod::Post, + AuthenticationMethod::Jwt, + ] { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let failure = json!({"error":"invalid_client","error_description":"client-authentication-secret"}); + let idp = if stage == Idp { + MockHttp::new(401, failure.clone()) + } else { + MockHttp::new(200, jag(&claims())) + }; + let resource = MockHttp::new(401, failure); + let error = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(method.configure("idp-client-secret", &provider)), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(method.configure("as-client-secret", &provider)), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange(&idp, &resource) + .await + .unwrap_err(); + assert_eq!( + error, + EmaError::OAuthRejected { + stage, + status: 401, + code: "invalid_client" + } + ); + assert!(!format!("{error:?} {error}").contains("client-authentication-secret")); + assert_eq!(idp.requests.lock().unwrap().len(), 1); + assert_eq!( + resource.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + let expected_calls = if method == AuthenticationMethod::Jwt { + 1 + usize::from(stage == ResourceServer) + } else { + 0 + }; + assert_eq!(provider.calls.lock().unwrap().len(), expected_calls); + } + } +} + +#[test] +fn client_authentication_debug_output_redacts_credentials_and_provider_state() { + const SECRET: &str = "client-authentication-secret-sentinel"; + let provider = RecordingAssertionProvider::new(AssertionBehavior::Empty(SECRET)); + let assertion = EmaClientAssertion::new(SECRET.into()); + assert!(!format!("{assertion:?}").contains(SECRET)); + for authentication in [ + EmaClientAuthentication::ClientSecretBasic(ClientSecret::new(SECRET.into())), + EmaClientAuthentication::ClientSecretPost(ClientSecret::new(SECRET.into())), + EmaClientAuthentication::JwtAssertion(provider.clone()), + ] { + assert!(!format!("{authentication:?}").contains(SECRET)); + let server = EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(authentication); + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new(server.clone(), server.clone(), RESOURCE, &refresh); + assert!(!format!("{server:?} {request:?}").contains(SECRET)); + assert!(!format!("{server:?} {request:?}").contains("private-query")); + } +} diff --git a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs index 2053920d2..165fccdb5 100644 --- a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs @@ -19,7 +19,10 @@ where /// 401 propagates as [`StreamableHttpError::AuthRequired`] carrying the /// `WWW-Authenticate` challenge for the caller to authorize with; /// - a token the server rejects (e.g. revoked) → one silent refresh, one - /// retry, then the challenge propagates. + /// retry, then the challenge propagates; + /// - a refresh that fails for any other reason (credential store, network, + /// provider) → that error propagates so the caller can retry instead of + /// being sent through a new authorization. async fn call_reacting_to_challenges( &self, auth_token: Option, @@ -54,10 +57,13 @@ where match refreshed { Ok(fresh_token) if fresh_token != sent_token => call(Some(fresh_token)).await, Ok(_) => Err(StreamableHttpError::AuthRequired(challenge)), - Err(error) => { - debug!("token refresh after server rejection failed: {error}"); + // `try_refresh_or_reauth` already reports the cases that need a + // new authorization; anything else is retryable or infrastructural. + Err(AuthError::AuthorizationRequired) => { + debug!("token refresh after server rejection requires authorization"); Err(StreamableHttpError::AuthRequired(challenge)) } + Err(error) => Err(error.into()), } } result => result, @@ -71,6 +77,10 @@ where { type Error = C::Error; + fn preserves_raw_responses() -> bool { + C::preserves_raw_responses() + } + async fn delete_session( &self, uri: std::sync::Arc, @@ -208,3 +218,167 @@ where .await } } + +#[cfg(all(test, feature = "transport-streamable-http-client-reqwest"))] +mod tests { + use std::sync::Arc; + + use oauth2::{AccessToken, RefreshToken, basic::BasicTokenType}; + + use super::*; + use crate::transport::{ + auth::{ + AuthorizationManager, AuthorizationMetadata, CredentialRefreshGuard, CredentialStore, + InMemoryCredentialStore, OAuthHttpClient, OAuthHttpClientFuture, OAuthHttpRequest, + OAuthTokenResponse, StoredCredentials, VendorExtraTokenFields, + }, + streamable_http_client::AuthRequiredError, + }; + + struct UnavailableStore; + + #[async_trait::async_trait] + impl CredentialStore for UnavailableStore { + async fn load(&self) -> Result, AuthError> { + unreachable!("guard failure must stop the credential load") + } + + async fn save(&self, _: StoredCredentials) -> Result<(), AuthError> { + unreachable!("guard failure must stop the credential save") + } + + async fn clear(&self) -> Result<(), AuthError> { + unreachable!("refresh must not clear credentials") + } + + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + Err(AuthError::CredentialStoreError("guard unavailable".into())) + } + } + + #[tokio::test] + async fn reactive_refresh_preserves_credential_store_failure() { + let mut manager = AuthorizationManager::new("https://mcp.example.com/mcp") + .await + .unwrap(); + manager.set_metadata(AuthorizationMetadata { + authorization_endpoint: "https://auth.example.com/authorize".into(), + token_endpoint: "https://auth.example.com/token".into(), + ..Default::default() + }); + manager.configure_client_id("client").unwrap(); + manager.set_credential_store(UnavailableStore); + let client = AuthClient::new(reqwest::Client::new(), manager); + + let error = client + .call_reacting_to_challenges(Some("old-token".into()), |_| async { + Err::<(), _>(StreamableHttpError::AuthRequired(AuthRequiredError::new( + "Bearer".into(), + ))) + }) + .await + .unwrap_err(); + + assert!(matches!(error, + StreamableHttpError::Auth(AuthError::CredentialStoreError(message)) + if message == "guard unavailable")); + } + + struct UnreachableTokenEndpoint; + + impl OAuthHttpClient for UnreachableTokenEndpoint { + fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + Box::pin(async { Err("token endpoint unreachable".into()) }) + } + } + + struct RejectingTokenEndpoint; + + impl OAuthHttpClient for RejectingTokenEndpoint { + fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + Box::pin(async { + Ok(oauth2::http::Response::builder() + .status(400) + .header("content-type", "application/json") + .body(br#"{"error":"invalid_grant"}"#.to_vec()) + .unwrap()) + }) + } + } + + /// A manager holding a refresh token the given token endpoint will answer for. + async fn manager_with_stored_refresh_token( + token_endpoint: Arc, + ) -> AuthorizationManager { + let mut manager = AuthorizationManager::new_with_oauth_http_client( + "https://mcp.example.com/mcp", + token_endpoint, + ) + .await + .unwrap(); + manager.set_metadata(AuthorizationMetadata { + authorization_endpoint: "https://auth.example.com/authorize".into(), + token_endpoint: "https://auth.example.com/token".into(), + ..Default::default() + }); + manager.configure_client_id("client").unwrap(); + + let mut token_response = OAuthTokenResponse::new( + AccessToken::new("old-token".into()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + token_response.set_refresh_token(Some(RefreshToken::new("stored-refresh".into()))); + let store = InMemoryCredentialStore::new(); + store + .save(StoredCredentials::new( + "client".into(), + Some(token_response), + vec![], + None, + )) + .await + .unwrap(); + manager.set_credential_store(store); + manager + } + + /// Drive one call whose server answer is a 401 challenge. + async fn challenge_once(manager: AuthorizationManager) -> StreamableHttpError { + AuthClient::new(reqwest::Client::new(), manager) + .call_reacting_to_challenges(Some("old-token".into()), |_| async { + Err::<(), _>(StreamableHttpError::AuthRequired(AuthRequiredError::new( + "Bearer".into(), + ))) + }) + .await + .unwrap_err() + } + + #[tokio::test] + async fn reactive_refresh_propagates_retryable_refresh_failure() { + let manager = manager_with_stored_refresh_token(Arc::new(UnreachableTokenEndpoint)).await; + + let error = challenge_once(manager).await; + + assert!( + matches!( + error, + StreamableHttpError::Auth(AuthError::TokenRefreshFailed(_)) + ), + "a retryable refresh failure must reach the caller instead of asking for a new authorization, got: {error:?}" + ); + } + + #[tokio::test] + async fn reactive_refresh_reports_a_rejected_refresh_token_as_a_challenge() { + let manager = manager_with_stored_refresh_token(Arc::new(RejectingTokenEndpoint)).await; + + let error = challenge_once(manager).await; + + assert!( + matches!(error, StreamableHttpError::AuthRequired(_)), + "a definitively rejected refresh token must surface the challenge, got: {error:?}" + ); + } +} diff --git a/crates/rmcp/src/transport/common/client_side_sse.rs b/crates/rmcp/src/transport/common/client_side_sse.rs index e0748ddbc..2be78ca15 100644 --- a/crates/rmcp/src/transport/common/client_side_sse.rs +++ b/crates/rmcp/src/transport/common/client_side_sse.rs @@ -205,6 +205,11 @@ impl Default for FixedInterval { pub struct ExponentialBackoff { pub max_times: Option, pub base_duration: Duration, + /// Optional upper bound on a single reconnect delay. `None` (the default) preserves the + /// pre-existing unbounded doubling behavior. Once the multiplier saturates near the bit + /// width that can still produce very long sleeps, so callers that need the client to + /// actually reconnect can set `Some(...)` to clamp the delay. + pub max_delay: Option, } impl ExponentialBackoff { @@ -216,6 +221,7 @@ impl Default for ExponentialBackoff { Self { max_times: None, base_duration: Self::DEFAULT_DURATION, + max_delay: None, } } } @@ -227,7 +233,16 @@ impl SseRetryPolicy for ExponentialBackoff { { return None; } - Some(self.base_duration * (2u32.pow(current_times as u32))) + // `current_times` is unbounded when `max_times` is unset, so the exponent can reach + // the bit width. Saturate the multiplier at `u32::MAX` and use saturating multiplication + // for the base duration so the delay stays monotonic and panic-free instead of an + // overflow panic (debug) or a wrapped-to-zero backoff (release). + let multiplier = 2u32.saturating_pow(current_times as u32); + let delay = self.base_duration.saturating_mul(multiplier); + Some(match self.max_delay { + Some(max_delay) => delay.min(max_delay), + None => delay, + }) } } @@ -776,4 +791,75 @@ mod tests { assert!(stream.next().await.is_none()); assert_eq!(attempts.load(Ordering::Relaxed), 0); } + + #[test] + fn exponential_backoff_saturates_at_high_retry_counts() { + // With `max_times` unset, `current_times` can reach the bit width. The old + // `2u32.pow(current_times)` panicked in debug builds and wrapped in release; + // the saturating implementation must return a monotonic, non-zero delay instead. + let policy = ExponentialBackoff { + max_times: None, + base_duration: Duration::from_millis(1), + max_delay: None, + }; + let mut previous = Duration::ZERO; + for current_times in [31usize, 32, 63, 64, 100] { + let delay = policy + .retry(current_times) + .expect("unbounded policy never gives up"); + assert!( + !delay.is_zero(), + "delay must stay non-zero at {current_times}" + ); + assert!( + delay >= previous, + "delay must stay monotonic at {current_times}" + ); + previous = delay; + } + } + + #[test] + fn exponential_backoff_caps_delay_at_max_delay() { + // An explicit cap keeps the unbounded doubling policy from producing decades-long + // sleeps once the multiplier saturates. The delay must grow monotonically, stop at + // the configured ceiling, and never exceed it. + let policy = ExponentialBackoff { + max_times: None, + base_duration: Duration::from_secs(1), + max_delay: Some(Duration::from_secs(30)), + }; + let mut previous = Duration::ZERO; + for current_times in [0usize, 1, 2, 3, 4, 5, 10, 32, 64, 100] { + let delay = policy + .retry(current_times) + .expect("unbounded policy never gives up"); + assert!( + delay >= previous, + "delay must stay monotonic at {current_times}" + ); + assert!( + delay <= Duration::from_secs(30), + "delay must respect max_delay at {current_times}" + ); + previous = delay; + } + // Beyond the ceiling the delay stays pinned at max_delay. + assert_eq!( + policy.retry(100).expect("never gives up"), + Duration::from_secs(30) + ); + } + + #[test] + fn exponential_backoff_respects_max_times() { + let policy = ExponentialBackoff { + max_times: Some(3), + base_duration: Duration::from_millis(1), + max_delay: None, + }; + assert!(policy.retry(0).is_some()); + assert!(policy.retry(2).is_some()); + assert!(policy.retry(3).is_none()); + } } diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index a8e59a244..e4e205a6d 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -114,6 +114,36 @@ fn cache_tools_from_response( } } +fn cache_tools_from_raw_response( + cache: &mut HashMap>, + message: &mut RawRxJsonRpcMessage, + protocol_version: &ProtocolVersion, +) { + if protocol_version < &ProtocolVersion::STANDARD_HEADERS { + return; + } + if let crate::model::JsonRpcMessage::Response(response) = message + && let Some(tools) = response + .result + .get_mut("tools") + .and_then(serde_json::Value::as_array_mut) + { + tools.retain(|value| { + // Preserve the original JSON, including extension fields. Malformed + // tools remain for the normal response decoder to reject. + let Ok(tool) = serde_json::from_value::(value.clone()) else { + return true; + }; + if let Err(reason) = mcp_headers::validate_param_header_annotations(&tool.input_schema) { + tracing::warn!(tool = %tool.name, "rejecting invalid x-mcp-header annotations: {reason}"); + return false; + } + cache.insert(tool.name.to_string(), tool.input_schema); + true + }); + } +} + fn negotiate_version_headers( init_response: &ServerJsonRpcMessage, base: HashMap, @@ -245,7 +275,7 @@ pub enum StreamableHttpProtocolError { MissingSessionIdInResponse, } -#[expect( +#[allow( clippy::large_enum_variant, reason = "boxing the streaming response would add an allocation to the common response path" )] @@ -390,7 +420,17 @@ pub(super) fn legacy_discover_response( /// handle the cancellation using state owned by that stream. pub trait StreamableHttpClient: Clone + Send + 'static { type Error: std::error::Error + Send + Sync + 'static; - /// Whether JSON responses are returned as [`StreamableHttpPostResponse::RawJson`]. + + /// Whether this backend preserves response result bodies as raw JSON. + /// + /// Returning `true` is a contract that all JSON responses use + /// [`StreamableHttpPostResponse::RawJson`] and that every SSE response path + /// yields events whose JSON-RPC result has not first been decoded through + /// [`ServerResult`]. Backends that cannot satisfy both requirements must + /// retain the default `false`; typed requests then fail before being sent. + /// The built-in reqwest and Unix-socket backends satisfy this contract. + /// The authenticated HTTP wrapper forwards the wrapped backend's + /// capability. fn preserves_raw_responses() -> bool { false } @@ -662,13 +702,7 @@ impl StreamableHttpClientWorker { } } - fn server_response_id( - message: &crate::model::JsonRpcMessage< - crate::model::ServerRequest, - Resp, - crate::model::ServerNotification, - >, - ) -> Option<&RequestId> { + fn server_response_id(message: &RawRxJsonRpcMessage) -> Option<&RequestId> { match message { crate::model::JsonRpcMessage::Response(response) => Some(&response.id), crate::model::JsonRpcMessage::Error(error) => error.id.as_ref(), @@ -676,13 +710,9 @@ impl StreamableHttpClientWorker { } } - fn clear_stream_response_pending( + fn clear_stream_response_pending( pending_stream_response_ids: &mut HashSet, - message: &crate::model::JsonRpcMessage< - crate::model::ServerRequest, - Resp, - crate::model::ServerNotification, - >, + message: &RawRxJsonRpcMessage, ) -> Option { let response_id = Self::server_response_id(message)?; if let Some(id) = pending_stream_response_ids.take(response_id) { @@ -856,24 +886,11 @@ impl StreamableHttpClientWorker { } } - async fn execute_sse_stream( + async fn execute_sse_stream( sse_stream: impl Stream< - Item = Result< - crate::model::JsonRpcMessage< - crate::model::ServerRequest, - Resp, - crate::model::ServerNotification, - >, - StreamableHttpError, - >, + Item = Result, StreamableHttpError>, > + Send, - sse_worker_tx: tokio::sync::mpsc::Sender< - crate::model::JsonRpcMessage< - crate::model::ServerRequest, - Resp, - crate::model::ServerNotification, - >, - >, + sse_worker_tx: tokio::sync::mpsc::Sender>, origin: InboundStreamOrigin, close_on_response: bool, ct: CancellationToken, @@ -1681,8 +1698,17 @@ impl Worker for StreamableHttpClientWorker { context.send_to_handler(message).await?; Ok(()) } - Ok(StreamableHttpPostResponse::RawJson(message, ..)) => { - context.send_to_handler(message).await?; + Ok(StreamableHttpPostResponse::RawJson(mut raw_message, ..)) => { + if matches!(&message, ClientJsonRpcMessage::Request(request) + if matches!(&request.request, ClientRequest::ListToolsRequest(_))) + { + cache_tools_from_raw_response( + &mut tool_header_cache, + &mut raw_message, + &version, + ); + } + context.send_to_handler(raw_message).await?; Ok(()) } Ok(StreamableHttpPostResponse::Sse(stream, ..)) => { @@ -1756,7 +1782,7 @@ impl Worker for StreamableHttpClientWorker { } let _ = responder.send(send_result); } - Event::ServerMessage(json_rpc_message) => { + Event::ServerMessage(mut json_rpc_message) => { // Match against all pending requests, not just open response streams. if let Some(request_id) = Self::clear_stream_response_pending( &mut pending_stream_response_ids, @@ -1764,15 +1790,11 @@ impl Worker for StreamableHttpClientWorker { ) { drop(request_stream_cancellations.remove(&request_id)); } - if let Ok(mut typed_message) = - decode_peer_response::(json_rpc_message.clone()) - { - cache_tools_from_response( - &mut tool_header_cache, - &mut typed_message, - &negotiated_version, - ); - } + cache_tools_from_raw_response( + &mut tool_header_cache, + &mut json_rpc_message, + &negotiated_version, + ); // send the message to the handler if let Err(e) = context.send_to_handler(json_rpc_message).await { break 'main_loop Err(e); @@ -2209,9 +2231,9 @@ mod tests { deprecated, reason = "Sampling is deprecated by SEP-2577 but remains the canonical restricted request" )] - fn sampling_request_message(id: i64) -> ServerJsonRpcMessage { + fn sampling_request_message(id: i64) -> RawRxJsonRpcMessage { use crate::model::{CreateMessageRequest, CreateMessageRequestParams, SamplingMessage}; - ServerJsonRpcMessage::request( + RawRxJsonRpcMessage::::request( ServerRequest::CreateMessageRequest(CreateMessageRequest::new( CreateMessageRequestParams::new(vec![SamplingMessage::user_text("hi")], 16), )), @@ -2225,8 +2247,9 @@ mod tests { InboundStreamOrigin::Unassociated, InboundStreamOrigin::OutboundRequest(RequestId::Number(3)), ] { - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .unwrap(), NumberOrString::Number(1), ); let stream = futures::stream::iter([Ok(sampling_request_message(9)), Ok(response)]); @@ -2241,7 +2264,7 @@ mod tests { .await .expect("stream completes"); - let ServerJsonRpcMessage::Request(request) = + let crate::model::JsonRpcMessage::Request(request) = rx.recv().await.expect("request forwarded") else { panic!("expected request first"); @@ -2254,7 +2277,7 @@ mod tests { // Responses are correlated by JSON-RPC id; no marker needed or added. assert!(matches!( rx.recv().await.expect("response forwarded"), - ServerJsonRpcMessage::Response(_) + crate::model::JsonRpcMessage::Response(_) )); } } @@ -2304,8 +2327,9 @@ mod tests { .lock() .expect("lock reconnects") .push((session_id.map(|id| id.to_string()), last_event_id)); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .expect("serialize response"), NumberOrString::Number(1), ); Ok(futures::stream::once(async move { @@ -2343,6 +2367,7 @@ mod tests { Arc::new(ExponentialBackoff { max_times: Some(1), base_duration: Duration::ZERO, + max_delay: None, }), ); let mut stream = std::pin::pin!(stream); @@ -2400,8 +2425,9 @@ mod tests { .expect("lock reconnects") .push((session_id.map(|id| id.to_string()), last_event_id)); let request = sampling_request_message(9); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .expect("serialize response"), NumberOrString::Number(1), ); // Stay open after the response, like a live connection, so the @@ -2447,6 +2473,7 @@ mod tests { Arc::new(ExponentialBackoff { max_times: Some(1), base_duration: Duration::ZERO, + max_delay: None, }), ); @@ -2531,6 +2558,36 @@ mod tests { ); } + #[test] + fn raw_tool_cache_preserves_extensions_and_filters_invalid_annotations() { + let mut valid = serde_json::to_value(tool( + "valid", + json!({"type": "string", "x-mcp-header": "Value"}), + )) + .unwrap(); + valid["vendorExtension"] = json!({"retained": true}); + let invalid = serde_json::to_value(tool( + "invalid", + json!({"type": "string", "x-mcp-header": ""}), + )) + .unwrap(); + let mut message = RawRxJsonRpcMessage::::response( + json!({"tools": [valid.clone(), invalid], "vendorResult": 42}), + NumberOrString::Number(1), + ); + let mut cache = HashMap::new(); + cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2026_07_28); + assert!(cache.contains_key("valid")); + assert!(!cache.contains_key("invalid")); + let crate::model::JsonRpcMessage::Response(response) = message else { + panic!("response") + }; + assert_eq!( + response.result, + json!({"tools": [valid], "vendorResult": 42}) + ); + } + #[test] fn cache_tools_preserves_pre_standard_header_results() { let invalid = tool("legacy", json!({ "type": "string", "x-mcp-header": "" })); @@ -2558,12 +2615,41 @@ mod tests { ); } + #[test] + fn raw_tool_cache_preserves_legacy_results_and_malformed_tools() { + let malformed = json!({"name": "broken", "inputSchema": "not-an-object"}); + let invalid = serde_json::to_value(tool("legacy", json!({"x-mcp-header": ""}))).unwrap(); + let original = json!({"tools": [invalid, malformed.clone()], "vendorResult": 42}); + let mut message = RawRxJsonRpcMessage::::response( + original.clone(), + NumberOrString::Number(1), + ); + let mut cache = HashMap::new(); + cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2025_11_25); + assert!(cache.is_empty()); + let crate::model::JsonRpcMessage::Response(response) = &message else { + panic!("response") + }; + assert_eq!(response.result, original); + + cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2026_07_28); + assert!(cache.is_empty()); + let crate::model::JsonRpcMessage::Response(response) = message else { + panic!("response") + }; + assert_eq!( + response.result, + json!({"tools": [malformed], "vendorResult": 42}) + ); + } + #[cfg(feature = "transport-streamable-http-client-reqwest")] #[test] fn clear_stream_response_pending_accepts_stringified_numeric_id() { let mut pending = HashSet::from([NumberOrString::Number(1)]); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .unwrap(), NumberOrString::String("1".into()), ); @@ -2582,8 +2668,9 @@ mod tests { fn clear_stream_response_pending_prefers_exact_string_id() { let string_id = NumberOrString::String("1".into()); let mut pending = HashSet::from([NumberOrString::Number(1), string_id.clone()]); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .unwrap(), string_id.clone(), ); diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index f1fef585f..59ef1a03c 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -55,6 +55,24 @@ use crate::{ pub(crate) const DEFAULT_MAX_REQUEST_BODY_BYTES: usize = 4 * 1024 * 1024; const STATELESS_STREAM_CHANNEL_CAPACITY: usize = 16; +struct ErrorResponse(Box); + +impl ErrorResponse { + fn into_response(self) -> BoxResponse { + *self.0 + } +} + +impl From for ErrorResponse { + fn from(response: BoxResponse) -> Self { + Self(Box::new(response)) + } +} + +type HttpResult = Result; +type RestoreResultSender = tokio::sync::watch::Sender>; +type PendingRestores = Arc>>; + #[non_exhaustive] #[derive(Debug, Clone)] pub struct StreamableHttpServerConfig { @@ -247,10 +265,6 @@ impl StreamableHttpServerConfig { } } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] /// Validates the `MCP-Protocol-Version` header on incoming HTTP requests. /// /// Per the MCP 2025-06-18 spec: @@ -259,7 +273,7 @@ impl StreamableHttpServerConfig { fn validate_protocol_version_header( headers: &http::HeaderMap, allow_unknown: bool, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { if let Some(value) = headers.get(HEADER_MCP_PROTOCOL_VERSION) { let version_str = value.to_str().map_err(|_| { Response::builder() @@ -284,7 +298,8 @@ fn validate_protocol_version_header( ))) .boxed(), ) - .expect("valid response")); + .expect("valid response") + .into()); } } Ok(()) @@ -322,7 +337,7 @@ impl> Service for NegotiatingStatelessHttpSer &requested, result.protocol_version.clone(), &self.0.supported_protocol_versions(), - ); + )?; if let Some(peer_info) = peer.peer_info() { let mut peer_info = (*peer_info).clone(); peer_info.protocol_version = result.protocol_version.clone(); @@ -349,48 +364,52 @@ impl> Service for NegotiatingStatelessHttpSer } } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] // SEP-2567: sessions are removed from the discover lifecycle. Validate // protocol-version consistency, then classify the request with the shared // lifecycle helper. fn is_legacy_request( message: Option<&ClientJsonRpcMessage>, headers: &HeaderMap, -) -> Result { +) -> HttpResult { let has_per_request_version = message.is_some_and(message_has_per_request_protocol_version); validate_protocol_version_header(headers, has_per_request_version)?; if let Some(message) = message { - if let ClientJsonRpcMessage::Request(req) = message { - if let ClientRequest::InitializeRequest(init) = &req.request { - validate_header_matches_init_body( - headers, - init.params.protocol_version.as_str(), - Some(req.id.clone()), - )?; - } + if let ClientJsonRpcMessage::Request(req) = message + && let ClientRequest::InitializeRequest(init) = &req.request + { + validate_header_matches_init_body( + headers, + init.params.protocol_version.as_str(), + Some(req.id.clone()), + )?; } validate_request_protocol_version_meta(headers, message)?; } + // An `initialize` request selects legacy semantics whatever version it names: + // the handshake exists only in the revisions before 2026-07-28, so the + // version in its params never routes it to the stateless path. The + // handshake itself answers with a legacy version the server supports. + if matches!( + message, + Some(ClientJsonRpcMessage::Request(req)) + if matches!(&req.request, ClientRequest::InitializeRequest(_)) + ) { + return Ok(true); + } + let uses_discover_lifecycle = matches!( message, Some(ClientJsonRpcMessage::Request(req)) - if !matches!(&req.request, ClientRequest::InitializeRequest(_)) - && req - .request - .get_meta() - .missing_required_keys(&ProtocolVersion::V_2026_07_28) - .is_empty() + if req + .request + .get_meta() + .missing_required_keys(&ProtocolVersion::V_2026_07_28) + .is_empty() ); let from_body = match message { - Some(ClientJsonRpcMessage::Request(req)) => match &req.request { - ClientRequest::InitializeRequest(init) => Some(init.params.protocol_version.clone()), - _ => req.request.get_meta().protocol_version(), - }, + Some(ClientJsonRpcMessage::Request(req)) => req.request.get_meta().protocol_version(), _ => None, }; let version = from_body @@ -422,10 +441,10 @@ async fn persist_and_forward_event( output: &mut Option>, ) -> Result<(), EventStoreError> { event.event_id = Some(event_store.store_event(stream_id, &event).await?); - if let Some(sender) = output { - if sender.send(event).await.is_err() { - *output = None; - } + if let Some(sender) = output + && sender.send(event).await.is_err() + { + *output = None; } Ok(()) } @@ -456,16 +475,12 @@ fn invalid_params_jsonrpc_response( .expect("valid response") } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] /// Absent header is allowed; the first initialize round-trip may legitimately omit it. fn validate_header_matches_init_body( headers: &http::HeaderMap, body_version: &str, request_id: Option, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let Some(header_value) = headers.get(HEADER_MCP_PROTOCOL_VERSION) else { return Ok(()); }; @@ -486,19 +501,16 @@ fn validate_header_matches_init_body( format!( "Invalid Request: MCP-Protocol-Version header ({header_str}) does not match initialize params.protocolVersion ({body_version})" ), - )); + ) + .into()); } Ok(()) } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_request_protocol_version_meta( headers: &HeaderMap, message: &ClientJsonRpcMessage, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let ClientJsonRpcMessage::Request(request) = message else { return Ok(()); }; @@ -522,7 +534,8 @@ fn validate_request_protocol_version_meta( "Invalid params: request _meta is missing or has malformed required fields: {}", missing.join(", ") ), - )); + ) + .into()); } return Ok(()); }; @@ -530,7 +543,8 @@ fn validate_request_protocol_version_meta( return Err(header_mismatch_jsonrpc_response( Some(request.id.clone()), "request _meta protocolVersion requires MCP-Protocol-Version header", - )); + ) + .into()); }; if header_version != meta_version.as_str() { return Err(header_mismatch_jsonrpc_response( @@ -538,7 +552,8 @@ fn validate_request_protocol_version_meta( format!( "MCP-Protocol-Version header ({header_version}) does not match request _meta protocolVersion ({meta_version})" ), - )); + ) + .into()); } Ok(()) } @@ -549,15 +564,11 @@ fn validate_request_protocol_version_meta( /// HTTP 400 / JSON-RPC `-32020` before handler dispatch. `server/discover` /// is included so the seam aligns with the per-POST header contract; its /// body-metadata rule is preserved unchanged. -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_required_protocol_header( config: &StreamableHttpServerConfig, headers: &HeaderMap, message: &ClientJsonRpcMessage, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { if !config.stateless_protocol_metadata_required { return Ok(()); } @@ -575,7 +586,8 @@ fn validate_required_protocol_header( Err(header_mismatch_jsonrpc_response( Some(request.id.clone()), "Missing MCP-Protocol-Version header for request requiring per-request protocol metadata", - )) + ) + .into()) } /// When `stateless_protocol_metadata_required` is enabled in stateless mode, @@ -585,14 +597,10 @@ fn validate_required_protocol_header( /// `server/discover` (whose body-metadata rule is already enforced by /// `validate_request_protocol_version_meta`), notifications, and other message /// kinds are exempt. -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_required_protocol_meta( config: &StreamableHttpServerConfig, message: &ClientJsonRpcMessage, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { if !config.stateless_protocol_metadata_required { return Ok(()); } @@ -611,7 +619,8 @@ fn validate_required_protocol_meta( Err(invalid_params_jsonrpc_response( Some(request.id.clone()), "Invalid params: request requires protocolVersion in request _meta", - )) + ) + .into()) } fn jsonrpc_http_status(message: &ServerJsonRpcMessage) -> http::StatusCode { @@ -632,7 +641,7 @@ fn jsonrpc_http_status(message: &ServerJsonRpcMessage) -> http::StatusCode { fn jsonrpc_message_response( message: ServerJsonRpcMessage, map_protocol_status: bool, -) -> Result { +) -> HttpResult { let status = if map_protocol_status { jsonrpc_http_status(&message) } else { @@ -666,15 +675,11 @@ fn header_mismatch_jsonrpc_response( /// The `initialize` handshake is exempt: clients emit these headers only after the /// version has been negotiated. `tool_schema` supplies the called tool's input schema /// so annotated `Mcp-Param-*` headers can be checked (no schema => those are skipped). -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_standard_headers( headers: &HeaderMap, message: &ClientJsonRpcMessage, tool_schema: impl Fn(&str) -> Option>, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let version_requires_headers = headers .get(HEADER_MCP_PROTOCOL_VERSION) .and_then(|value| value.to_str().ok()) @@ -707,7 +712,7 @@ fn validate_standard_headers( .and_then(|name| name.as_str()) .and_then(tool_schema); if let Err(reason) = mcp_headers::validate_request_headers(headers, &value, schema.as_deref()) { - return Err(header_mismatch_jsonrpc_response(request_id, reason)); + return Err(header_mismatch_jsonrpc_response(request_id, reason).into()); } Ok(()) } @@ -786,13 +791,22 @@ fn parse_origin_value(value: &str) -> Option { if value.eq_ignore_ascii_case("null") { return Some(NormalizedOrigin::Null); } + let (_, serialized_authority) = value.split_once("://")?; + if serialized_authority.contains(['/', '?', '#', '@']) { + return None; + } let uri = http::Uri::try_from(value).ok()?; let scheme = uri.scheme_str()?.to_ascii_lowercase(); let authority = uri.authority()?; + let port = authority.port_u16().or(match scheme.as_str() { + "http" => Some(80), + "https" => Some(443), + _ => None, + }); Some(NormalizedOrigin::Tuple { scheme, host: normalize_host(authority.host()), - port: authority.port_u16(), + port, }) } @@ -816,7 +830,7 @@ fn origin_is_allowed(origin: &NormalizedOrigin, allowed_origins: &[String]) -> b host: o_host, port: o_port, }, - ) => a_scheme == o_scheme && a_host == o_host && (a_port.is_none() || a_port == o_port), + ) => a_scheme == o_scheme && a_host == o_host && a_port == o_port, _ => false, }) } @@ -831,10 +845,7 @@ fn bad_request_response(message: &str) -> BoxResponse { .expect("failed to build bad request response") } -fn parse_host_header( - uri: &http::Uri, - headers: &HeaderMap, -) -> Result { +fn parse_host_header(uri: &http::Uri, headers: &HeaderMap) -> HttpResult { if let Some(host) = headers.get(http::header::HOST) { let host_str = host .to_str() @@ -865,50 +876,50 @@ fn validate_dns_rebinding_headers( uri: &http::Uri, headers: &HeaderMap, config: &StreamableHttpServerConfig, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let host = parse_host_header(uri, headers)?; if !host_is_allowed(&host, &config.allowed_hosts) { tracing::warn!( host = ?host, "rejected request with disallowed Host header (possible DNS rebinding attempt)", ); - return Err(forbidden_response("Forbidden: Host header is not allowed")); + return Err(forbidden_response("Forbidden: Host header is not allowed").into()); } validate_origin_header(headers, &config.allowed_origins)?; Ok(()) } -fn validate_origin_header( - headers: &HeaderMap, - allowed_origins: &[String], -) -> Result<(), BoxResponse> { +fn validate_origin_header(headers: &HeaderMap, allowed_origins: &[String]) -> HttpResult<()> { if allowed_origins.is_empty() { return Ok(()); } - let Some(origin_header) = headers.get(http::header::ORIGIN) else { + let mut origin_headers = headers.get_all(http::header::ORIGIN).iter(); + let Some(origin_header) = origin_headers.next() else { return Ok(()); }; + if origin_headers.next().is_some() { + tracing::warn!("rejected request with multiple Origin headers"); + return Err(forbidden_response("Forbidden: Multiple Origin headers").into()); + } let origin_str = origin_header .to_str() .inspect_err(|_| { tracing::warn!(origin = ?origin_header, "rejected request with non-UTF-8 Origin header"); }) - .map_err(|_| bad_request_response("Bad Request: Invalid Origin header encoding"))?; + .map_err(|_| forbidden_response("Forbidden: Invalid Origin header encoding"))?; let origin = parse_origin_value(origin_str).ok_or_else(|| { tracing::warn!( origin = origin_str, "rejected request with malformed Origin header", ); - bad_request_response("Bad Request: Invalid Origin header") + forbidden_response("Forbidden: Invalid Origin header") })?; if !origin_is_allowed(&origin, allowed_origins) { tracing::warn!( origin = ?origin, "rejected request with disallowed Origin header (possible cross-origin attack)", ); - return Err(forbidden_response( - "Forbidden: Origin header is not allowed", - )); + return Err(forbidden_response("Forbidden: Origin header is not allowed").into()); } Ok(()) } @@ -1004,9 +1015,7 @@ pub struct StreamableHttpService { /// same unknown session ID wait for the first restore to complete rather /// than racing to replay the initialize handshake. `None` when no external /// session store is configured (avoids allocating the map). - pending_restores: Option< - Arc>>>>, - >, + pending_restores: Option, /// Caches tool input schemas by name for SEP-2243 `Mcp-Param-*` validation. /// Populated lazily via `get_tool` so the service factory runs at most once /// per tool name. `None` value means the tool exposes no schema. @@ -1059,10 +1068,9 @@ where /// `result` defaults to `false` (failure / cancellation). Only the success path /// needs to set it to `true` before returning. struct PendingRestoreGuard { - pending_restores: - Arc>>>>, + pending_restores: PendingRestores, session_id: SessionId, - watch_tx: tokio::sync::watch::Sender>, + watch_tx: RestoreResultSender, /// The value that will be broadcast to waiting tasks on drop. result: bool, } @@ -1090,12 +1098,10 @@ where session_manager: Arc, config: StreamableHttpServerConfig, ) -> Self { - let pending_restores = config.session_store.is_some().then(|| { - Arc::new(tokio::sync::RwLock::new(HashMap::< - SessionId, - tokio::sync::watch::Sender>, - >::new())) - }); + let pending_restores = config + .session_store + .is_some() + .then(|| Arc::new(tokio::sync::RwLock::new(HashMap::new()))); Self { config, session_manager, @@ -1122,19 +1128,18 @@ where tokio::spawn(async move { let mut sender = Some(sender); - if let Some(retry) = retry { - if let Err(error) = persist_and_forward_event( + if let Some(retry) = retry + && let Err(error) = persist_and_forward_event( event_store.as_ref(), &stream_id, ServerSseMessage::retry(retry), &mut sender, ) .await - { - tracing::error!(%stream_id, %error, "failed to persist SSE priming event"); - request_ct.cancel(); - return; - } + { + tracing::error!(%stream_id, %error, "failed to persist SSE priming event"); + request_ct.cancel(); + return; } let mut first = first; @@ -1206,7 +1211,7 @@ where service: S, mut request: crate::model::JsonRpcRequest, parts: http::request::Parts, - ) -> Result { + ) -> HttpResult { let peer_info = Self::peer_info_for_stateless_request(&request, &parts.headers); request.request.extensions_mut().insert(parts); let (transport, mut receiver) = @@ -1271,10 +1276,10 @@ where /// per name to read its `ServerHandler::get_tool` definition. Used to /// validate SEP-2243 `Mcp-Param-*` headers against the request body. fn tool_schema(&self, name: &str) -> Option> { - if let Ok(cache) = self.tool_schemas.read() { - if let Some(schema) = cache.get(name) { - return schema.clone(); - } + if let Ok(cache) = self.tool_schemas.read() + && let Some(schema) = cache.get(name) + { + return schema.clone(); } let schema = self .get_service() @@ -1464,23 +1469,15 @@ where Some(init_done_tx), ); - if let Err(e) = self - .session_manager + self.session_manager .initialize_session(session_id, restore_init) .await - .map_err(|e| std::io::Error::other(e.to_string())) - { - return Err(e); - } + .map_err(|e| std::io::Error::other(e.to_string()))?; - if let Err(e) = self - .session_manager + self.session_manager .accept_message(session_id, restore_initialized) .await - .map_err(|e| std::io::Error::other(e.to_string())) - { - return Err(e); - } + .map_err(|e| std::io::Error::other(e.to_string()))?; if init_done_rx.await.is_err() { return Err(std::io::Error::other( @@ -1505,7 +1502,7 @@ where if let Err(response) = validate_dns_rebinding_headers(request.uri(), request.headers(), &self.config) { - return response; + return response.into_response(); } let method = request.method().clone(); let supports_stateless_replay = self.session_manager.event_store().is_some(); @@ -1532,10 +1529,10 @@ where }; match result { Ok(response) => response, - Err(response) => response, + Err(response) => response.into_response(), } } - async fn handle_get(&self, request: Request) -> Result + async fn handle_get(&self, request: Request) -> HttpResult where B: Body + Send + 'static, B::Error: Display, @@ -1678,7 +1675,7 @@ where )) } - async fn handle_post(&self, request: Request) -> Result + async fn handle_post(&self, request: Request) -> HttpResult where B: Body + Send + 'static, B::Error: Display, @@ -1830,7 +1827,7 @@ where let stored_init_params = match &mut message { ClientJsonRpcMessage::Request(req) => { let ClientRequest::InitializeRequest(init_req) = &req.request else { - return Err(unexpected_message_response("initialize request")); + return Err(unexpected_message_response("initialize request").into()); }; // Reject mismatched MCP-Protocol-Version header before binding the session to anything. validate_header_matches_init_body( @@ -1848,7 +1845,7 @@ where stored_init_params } _ => { - return Err(unexpected_message_response("initialize request")); + return Err(unexpected_message_response("initialize request").into()); } }; let service = self @@ -2005,7 +2002,8 @@ where std::io::ErrorKind::UnexpectedEof, "no response message received from handler", ), - )); + ) + .into()); }; tracing::trace!(?message); if matches!( @@ -2037,7 +2035,7 @@ where } } - async fn handle_delete(&self, request: Request) -> Result + async fn handle_delete(&self, request: Request) -> HttpResult where B: Body + Send + 'static, B::Error: Display, diff --git a/crates/rmcp/src/transport/worker.rs b/crates/rmcp/src/transport/worker.rs index 6f2256f07..8924ecad0 100644 --- a/crates/rmcp/src/transport/worker.rs +++ b/crates/rmcp/src/transport/worker.rs @@ -68,6 +68,9 @@ pub trait Worker: Sized + Send + 'static { fn config(&self) -> WorkerConfig { WorkerConfig::default() } + fn preserves_raw_responses() -> bool { + false + } /// Return true to send this message through the separate control queue. /// /// Workers that opt in must read [`WorkerContext::control_from_handler_rx`] @@ -81,10 +84,6 @@ pub trait Worker: Sized + Send + 'static { fn supports_request_cancellation() -> bool { false } - /// Whether this worker sends inbound responses through the lossless raw path. - fn preserves_raw_responses() -> bool { - false - } } type RequestCancellations = Arc>>>; @@ -333,12 +332,12 @@ pub struct SendRequest { #[non_exhaustive] pub struct WorkerContext { pub to_handler_tx: tokio::sync::mpsc::Sender>, + raw_to_handler_tx: tokio::sync::mpsc::Sender>, pub from_handler_rx: tokio::sync::mpsc::Receiver>, /// Messages selected by [`Worker::is_control_message`]. pub control_from_handler_rx: tokio::sync::mpsc::Receiver>, pub cancellation_token: CancellationToken, control_generation: Arc, - raw_to_handler_tx: tokio::sync::mpsc::Sender>, } impl WorkerContext { @@ -373,23 +372,19 @@ impl WorkerContext { Resp: serde::Serialize, { let item = match item { - crate::model::JsonRpcMessage::Request(request) => { - crate::model::JsonRpcMessage::Request(request) - } - crate::model::JsonRpcMessage::Response(response) => { - crate::model::JsonRpcMessage::Response(crate::model::JsonRpcResponse { + JsonRpcMessage::Request(request) => JsonRpcMessage::Request(request), + JsonRpcMessage::Response(response) => { + JsonRpcMessage::Response(crate::model::JsonRpcResponse { jsonrpc: response.jsonrpc, id: response.id, result: serde_json::to_value(response.result) .map_err(WorkerQuitReason::ResponseSerialization)?, }) } - crate::model::JsonRpcMessage::Notification(notification) => { - crate::model::JsonRpcMessage::Notification(notification) - } - crate::model::JsonRpcMessage::Error(error) => { - crate::model::JsonRpcMessage::Error(error) + JsonRpcMessage::Notification(notification) => { + JsonRpcMessage::Notification(notification) } + JsonRpcMessage::Error(error) => JsonRpcMessage::Error(error), }; self.raw_to_handler_tx .send(item) @@ -469,8 +464,7 @@ impl Transport for WorkerTransport { } async fn receive(&mut self) -> Option> { loop { - let message = self.rx.recv().await?; - match decode_peer_response::(message) { + match decode_peer_response::(self.rx.recv().await?) { Ok(message) => return Some(message), Err(error) => { tracing::debug!(%error, "Ignoring response with invalid result shape") diff --git a/crates/rmcp/tests/support/typed_stdio_server.sh b/crates/rmcp/tests/support/typed_stdio_server.sh new file mode 100644 index 000000000..85850380d --- /dev/null +++ b/crates/rmcp/tests/support/typed_stdio_server.sh @@ -0,0 +1,20 @@ +#!/bin/sh + +# Minimal line-delimited MCP server used to exercise TokioChildProcess without +# relying on a language-specific MCP implementation or package installation. +while IFS= read -r message; do + request_id=$(printf '%s\n' "$message" | sed -n 's/.*"id":\([^,}]*\).*/\1/p') + case "$message" in + *'"method":"initialize"'*) + printf '%s\n' "{\"jsonrpc\":\"2.0\",\"id\":${request_id},\"result\":{\"protocolVersion\":\"2025-03-26\",\"capabilities\":{},\"serverInfo\":{\"name\":\"typed-stdio-fixture\",\"version\":\"1.0.0\"}}}" + ;; + *'"method":"skills/list"'*) + printf '%s\n' "{\"jsonrpc\":\"2.0\",\"id\":${request_id},\"result\":{\"resultType\":\"complete\",\"skills\":[\"stdio-example\"],\"_meta\":{\"vendorExtension\":{\"retained\":true}}}}" + ;; + *'"method":"ping"'*) + printf '%s\n' "{\"jsonrpc\":\"2.0\",\"id\":${request_id},\"result\":{}}" + ;; + *'"method":"notifications/initialized"'*) + ;; + esac +done diff --git a/crates/rmcp/tests/test_custom_headers.rs b/crates/rmcp/tests/test_custom_headers.rs index 736dce18e..250f97dd4 100644 --- a/crates/rmcp/tests/test_custom_headers.rs +++ b/crates/rmcp/tests/test_custom_headers.rs @@ -380,10 +380,10 @@ async fn test_mcp_custom_headers_sent_to_server() -> anyhow::Result<()> { let mut headers_map = HashMap::new(); for (name, value) in headers.iter() { let name_str = name.as_str(); - if name_str.starts_with("x-") { - if let Ok(v) = value.to_str() { - headers_map.insert(name_str.to_string(), v.to_string()); - } + if name_str.starts_with("x-") + && let Ok(v) = value.to_str() + { + headers_map.insert(name_str.to_string(), v.to_string()); } } @@ -392,48 +392,48 @@ async fn test_mcp_custom_headers_sent_to_server() -> anyhow::Result<()> { stored.extend(headers_map); // Parse the MCP request - if let Ok(json_body) = serde_json::from_slice::(&body) { - if let Some(method) = json_body.get("method").and_then(|m| m.as_str()) { - if method == "initialize" { - state.initialize_called.notify_one(); - // Return a valid MCP initialize response with session header - let response = json!({ - "jsonrpc": "2.0", - "id": json_body.get("id"), - "result": { - "protocolVersion": "2024-11-05", - "capabilities": {}, - "serverInfo": { - "name": "test-server", - "version": "1.0.0" - } + if let Ok(json_body) = serde_json::from_slice::(&body) + && let Some(method) = json_body.get("method").and_then(|m| m.as_str()) + { + if method == "initialize" { + state.initialize_called.notify_one(); + // Return a valid MCP initialize response with session header + let response = json!({ + "jsonrpc": "2.0", + "id": json_body.get("id"), + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "serverInfo": { + "name": "test-server", + "version": "1.0.0" } - }); - return ( - StatusCode::OK, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "test-session-123", - ), - ], - response.to_string(), - ); - } else if method == "notifications/initialized" { - // For initialized notification, return 202 Accepted - return ( - StatusCode::ACCEPTED, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "test-session-123", - ), - ], - String::new(), - ); - } + } + }); + return ( + StatusCode::OK, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "test-session-123", + ), + ], + response.to_string(), + ); + } else if method == "notifications/initialized" { + // For initialized notification, return 202 Accepted + return ( + StatusCode::ACCEPTED, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "test-session-123", + ), + ], + String::new(), + ); } } @@ -1124,7 +1124,10 @@ async fn test_server_falls_back_to_uri_authority_when_host_header_missing() { #[cfg(all(feature = "transport-streamable-http-server", feature = "server"))] mod origin_validation { - use std::sync::Arc; + use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }; use bytes::Bytes; use http::{Method, Request, header::CONTENT_TYPE}; @@ -1149,9 +1152,19 @@ mod origin_validation { fn service_with_allowed_origins( origins: &[&str], + ) -> StreamableHttpService { + service_with_allowed_origins_and_counter(origins, Arc::new(AtomicUsize::new(0))) + } + + fn service_with_allowed_origins_and_counter( + origins: &[&str], + handler_creations: Arc, ) -> StreamableHttpService { StreamableHttpService::new( - || Ok(TestHandler), + move || { + handler_creations.fetch_add(1, Ordering::SeqCst); + Ok(TestHandler) + }, Arc::new(LocalSessionManager::default()), StreamableHttpServerConfig::default().with_allowed_origins(origins.iter().copied()), ) @@ -1199,6 +1212,80 @@ mod origin_validation { assert_eq!(response.status(), http::StatusCode::FORBIDDEN); } + #[tokio::test] + async fn malformed_origin_is_forbidden() { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + let response = service.handle(init_request(Some("not an origin"))).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn non_utf8_origin_is_forbidden_before_handler_creation() { + let handler_creations = Arc::new(AtomicUsize::new(0)); + let service = service_with_allowed_origins_and_counter( + &["http://localhost:8080"], + handler_creations.clone(), + ); + let mut request = init_request(None); + request.headers_mut().insert( + http::header::ORIGIN, + http::HeaderValue::from_bytes(b"\xff").expect("opaque header value"), + ); + + let response = service.handle(request).await; + + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + assert_eq!(handler_creations.load(Ordering::SeqCst), 0); + } + + #[tokio::test] + async fn multiple_origin_headers_are_forbidden() { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + let mut request = init_request(Some("http://localhost:8080")); + request.headers_mut().append( + http::header::ORIGIN, + "http://attacker.example".parse().unwrap(), + ); + let response = service.handle(request).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn allowlisted_origin_does_not_wildcard_an_unexpected_port() { + let service = service_with_allowed_origins(&["https://app.example"]); + let response = service + .handle(init_request(Some("https://app.example:9443"))) + .await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn default_origin_ports_are_equivalent() { + for (allowed, presented) in [ + ("http://localhost", "http://localhost:80"), + ("https://localhost:443", "https://localhost"), + ] { + let service = service_with_allowed_origins(&[allowed]); + let response = service.handle(init_request(Some(presented))).await; + assert_eq!(response.status(), http::StatusCode::OK); + } + } + + #[tokio::test] + async fn origin_with_non_origin_components_is_forbidden() { + for origin in [ + "http://localhost:8080/", + "http://localhost:8080/evil", + "http://localhost:8080?query=1", + "http://user@localhost:8080", + "http://localhost:8080#fragment", + ] { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + let response = service.handle(init_request(Some(origin))).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN, "{origin}"); + } + } + #[tokio::test] async fn missing_origin_passes_through() { let service = service_with_allowed_origins(&["http://localhost:8080"]); diff --git a/crates/rmcp/tests/test_custom_request.rs b/crates/rmcp/tests/test_custom_request.rs index 6ccc5992f..b4056b5ff 100644 --- a/crates/rmcp/tests/test_custom_request.rs +++ b/crates/rmcp/tests/test_custom_request.rs @@ -1,5 +1,12 @@ #![cfg(not(feature = "local"))] -use std::{future::Future, sync::Arc, time::Duration}; +use std::{ + future::Future, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; use rmcp::{ ClientHandler, RoleClient, ServerHandler, ServiceExt, @@ -7,7 +14,7 @@ use rmcp::{ ClientRequest, ClientResult, CustomRequest, CustomResult, ErrorCode, ErrorData, PingRequest, ServerRequest, ServerResult, }, - service::{PeerRequestOptions, ServiceError}, + service::{PeerRequestOptions, ServiceError, serve_directly}, transport::Transport, }; use serde::Deserialize; @@ -44,6 +51,83 @@ async fn existing_transport_implementations_get_the_raw_receive_compatibility_de assert!(transport.receive_raw().await.is_none()); } +struct LegacyResponseTransport { + responses: tokio::sync::mpsc::UnboundedReceiver, + response_tx: tokio::sync::mpsc::UnboundedSender, + sends: Arc, +} + +impl Transport for LegacyResponseTransport { + type Error = std::convert::Infallible; + + fn send( + &mut self, + item: rmcp::service::TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + let response_tx = self.response_tx.clone(); + let sends = self.sends.clone(); + async move { + if let rmcp::model::ClientJsonRpcMessage::Request(request) = item { + sends.fetch_add(1, Ordering::SeqCst); + let _ = response_tx.send(rmcp::model::ServerJsonRpcMessage::response( + ServerResult::CustomResult(CustomResult::new(json!({}))), + request.id, + )); + } + Ok(()) + } + } + + async fn receive(&mut self) -> Option> { + self.responses.recv().await + } + + fn close(&mut self) -> impl Future> + Send { + std::future::ready(Ok(())) + } +} + +#[tokio::test] +async fn legacy_transport_preserves_standard_response_without_raw_round_trip() -> anyhow::Result<()> +{ + let (response_tx, responses) = tokio::sync::mpsc::unbounded_channel(); + let sends = Arc::new(AtomicUsize::new(0)); + let transport = LegacyResponseTransport { + responses, + response_tx, + sends: sends.clone(), + }; + let client = serve_directly::( + (), + transport, + Some(rmcp::model::ServerInfo::default().into()), + ); + + let typed = client + .send_request_as::(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await; + assert!(matches!(typed, Err(ServiceError::RawResponseUnavailable))); + assert_eq!(sends.load(Ordering::SeqCst), 0, "typed request was sent"); + + let standard = client + .send_request(ClientRequest::CustomRequest(CustomRequest::new( + "requests/custom-test", + None, + ))) + .await?; + assert!(matches!( + standard, + ServerResult::CustomResult(result) if result.0 == json!({}) + )); + assert_eq!(sends.load(Ordering::SeqCst), 1); + + client.cancel().await?; + Ok(()) +} + struct CustomRequestServer { receive_signal: Arc, payload: Arc>>, @@ -170,6 +254,217 @@ async fn typed_test_client() -> anyhow::Result anyhow::Result<()> { + use rmcp::transport::{ + auth::{AuthClient, AuthorizationManager}, + streamable_http_client::{StreamableHttpClientTransportConfig, StreamableHttpClientWorker}, + streamable_http_server::{ + StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager, + }, + }; + let stop = tokio_util::sync::CancellationToken::new(); + let server: StreamableHttpService = + StreamableHttpService::new( + || Ok(TypedCustomRequestServer), + Default::default(), + StreamableHttpServerConfig::default().with_cancellation_token(stop.child_token()), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?; + let endpoint = format!("http://{}/mcp", listener.local_addr()?); + let router = axum::Router::new().nest_service("/mcp", server); + let shutdown = stop.clone(); + let task = tokio::spawn(async move { + axum::serve(listener, router) + .with_graceful_shutdown(shutdown.cancelled_owned()) + .await + }); + let result = tokio::time::timeout(Duration::from_secs(10), async { + let manager = AuthorizationManager::new(&endpoint).await?; + let auth = AuthClient::new(reqwest::Client::new(), manager); + let transport = StreamableHttpClientWorker::new( + auth, + StreamableHttpClientTransportConfig::with_uri(endpoint), + ); + let client = ().serve(transport).await?; + let response = client + .send_request_as::(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + Some(json!({})), + ))) + .await; + client.cancel().await?; + let response = response?; + assert_eq!(response.skills, ["example"]); + assert_eq!( + response.meta["io.modelcontextprotocol/serverInfo"]["name"], + "skills" + ); + anyhow::Ok(()) + }) + .await; + stop.cancel(); + tokio::time::timeout(Duration::from_secs(5), task).await???; + result??; + Ok(()) +} + +#[cfg(all( + feature = "auth", + feature = "transport-streamable-http-server", + feature = "transport-streamable-http-client-reqwest" +))] +#[tokio::test] +async fn json_tool_listing_populates_parameter_headers_through_auth_client() -> anyhow::Result<()> { + use rmcp::{ + ClientServiceExt, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, ListToolsResult, + PaginatedRequestParams, ServerCapabilities, ServerInfo, Tool, + }, + service::RequestContext, + transport::{ + auth::{AuthClient, AuthorizationManager}, + streamable_http_client::{ + StreamableHttpClientTransportConfig, StreamableHttpClientWorker, + }, + streamable_http_server::{ + StreamableHttpServerConfig, StreamableHttpService, + session::local::LocalSessionManager, + }, + }, + }; + + struct AnnotatedToolServer; + impl ServerHandler for AnnotatedToolServer { + fn supported_protocol_versions( + &self, + ) -> std::borrow::Cow<'static, [rmcp::model::ProtocolVersion]> { + std::borrow::Cow::Borrowed(&[rmcp::model::ProtocolVersion::V_2026_07_28]) + } + + fn get_info(&self) -> ServerInfo { + ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) + } + + fn get_tool(&self, name: &str) -> Option { + (name == "deploy").then(|| { + Tool::new( + "deploy", + "Deploy in a region", + Arc::new( + json!({"type": "object", "properties": { + "region": {"type": "string", "x-mcp-header": "Region"} + }}) + .as_object() + .unwrap() + .clone(), + ), + ) + }) + } + + async fn list_tools( + &self, + _: Option, + _: RequestContext, + ) -> Result { + Ok(ListToolsResult::with_all_items(vec![ + self.get_tool("deploy").unwrap(), + ])) + } + + async fn call_tool( + &self, + _: CallToolRequestParams, + _: RequestContext, + ) -> Result { + Ok(CallToolResult::success(vec![]).into()) + } + } + + let observed_headers = Arc::new(Mutex::new(Vec::new())); + let capture = observed_headers.clone(); + let stop = tokio_util::sync::CancellationToken::new(); + let server: StreamableHttpService = + StreamableHttpService::new( + || Ok(AnnotatedToolServer), + Default::default(), + StreamableHttpServerConfig::default() + .with_legacy_session_mode(false) + .with_json_response(true) + .with_cancellation_token(stop.child_token()), + ); + let router = axum::Router::new() + .nest_service("/mcp", server) + .layer(axum::middleware::from_fn( + move |request: axum::extract::Request, next: axum::middleware::Next| { + let capture = capture.clone(); + async move { + if request + .headers() + .get("Mcp-Method") + .is_some_and(|method| method == "tools/call") + { + capture + .lock() + .await + .push(request.headers().get("Mcp-Param-Region").cloned()); + } + next.run(request).await + } + }, + )); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?; + let endpoint = format!("http://{}/mcp", listener.local_addr()?); + let shutdown = stop.clone(); + let task = tokio::spawn(async move { + axum::serve(listener, router) + .with_graceful_shutdown(shutdown.cancelled_owned()) + .await + }); + let result = tokio::time::timeout(Duration::from_secs(10), async { + let manager = AuthorizationManager::new(&endpoint).await?; + let transport = StreamableHttpClientWorker::new( + AuthClient::new(reqwest::Client::new(), manager), + StreamableHttpClientTransportConfig::with_uri(endpoint), + ); + let client = () + .serve_with_lifecycle( + transport, + rmcp::service::ClientLifecycleMode::Discover { + preferred_versions: vec![rmcp::model::ProtocolVersion::V_2026_07_28], + }, + ) + .await?; + let listed = client.list_tools(None).await?; + assert_eq!(listed.tools.len(), 1); + let response = client + .call_tool( + CallToolRequestParams::new("deploy") + .with_arguments(json!({"region": "us-east-1"}).as_object().unwrap().clone()), + ) + .await; + client.cancel().await?; + response?; + assert_eq!( + *observed_headers.lock().await, + vec![Some(http::HeaderValue::from_static("us-east-1"))] + ); + anyhow::Ok(()) + }) + .await; + stop.cancel(); + tokio::time::timeout(Duration::from_secs(5), task).await???; + result??; + Ok(()) +} + #[tokio::test] async fn typed_custom_request_bypasses_the_core_response_union() -> anyhow::Result<()> { let client = typed_test_client().await?; @@ -252,6 +547,73 @@ async fn typed_custom_requests_use_standard_timeout_and_cancellation() -> anyhow Ok(()) } +struct CancellationObservingServer { + started: Arc, + cancelled: Arc, +} + +impl ServerHandler for CancellationObservingServer { + async fn on_custom_request( + &self, + request: CustomRequest, + context: rmcp::service::RequestContext, + ) -> Result { + if request.method == "skills/list" { + return Ok(CustomResult::new(json!({ + "resultType": "complete", + "skills": ["after-cancellation"], + "_meta": {} + }))); + } + self.started.notify_one(); + context.ct.cancelled().await; + self.cancelled.notify_one(); + Ok(CustomResult::new(json!({"cancelled": true}))) + } +} + +#[tokio::test] +async fn typed_custom_request_handle_sends_cancellation_to_the_peer() -> anyhow::Result<()> { + let started = Arc::new(Notify::new()); + let cancelled = Arc::new(Notify::new()); + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_task = tokio::spawn({ + let started = started.clone(); + let cancelled = cancelled.clone(); + async move { + CancellationObservingServer { started, cancelled } + .serve(server_transport) + .await? + .waiting() + .await?; + anyhow::Ok(()) + } + }); + let client = ().serve(client_transport).await?; + let handle = client + .send_request_as_with_option::( + ClientRequest::CustomRequest(CustomRequest::new("skills/cancel", None)), + PeerRequestOptions::no_options(), + ) + .await?; + + tokio::time::timeout(Duration::from_secs(2), started.notified()).await?; + handle.cancel(Some("test cancellation".into())).await?; + tokio::time::timeout(Duration::from_secs(2), cancelled.notified()).await?; + + let following: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await?; + assert_eq!(following.skills, ["after-cancellation"]); + + client.cancel().await?; + tokio::time::timeout(Duration::from_secs(2), server_task).await???; + Ok(()) +} + struct CustomRequestClient { receive_signal: Arc, payload: Arc>>, diff --git a/crates/rmcp/tests/test_handler_cache_hints.rs b/crates/rmcp/tests/test_handler_cache_hints.rs index d6b626d2a..7f44e8280 100644 --- a/crates/rmcp/tests/test_handler_cache_hints.rs +++ b/crates/rmcp/tests/test_handler_cache_hints.rs @@ -2,10 +2,14 @@ #![cfg(feature = "client")] use rmcp::{ - ClientHandler, ServerHandler, ServiceExt, + ClientHandler, RoleClient, RoleServer, ServerHandler, handler::server::router::{prompt::PromptRouter, tool::ToolRouter}, - model::{CacheScope, ClientInfo, ListPromptsResult, ListToolsResult, ProtocolVersion}, - prompt_handler, tool_handler, + model::{ + CacheScope, ClientInfo, ListPromptsResult, ListToolsResult, ProtocolVersion, ServerInfo, + }, + prompt_handler, + service::serve_directly, + tool_handler, }; #[derive(Debug, Clone)] @@ -40,22 +44,33 @@ impl ClientHandler for VersionedClient { } } +/// Wires the pair up directly on `protocol_version`. `2026-07-28` removed the +/// `initialize` handshake, so a peer on that revision is reached the way the +/// discover lifecycle leaves one: with the version already agreed. async fn list_results(protocol_version: ProtocolVersion) -> (ListToolsResult, ListPromptsResult) { let (server_transport, client_transport) = tokio::io::duplex(4096); + let client_handler = VersionedClient { + protocol_version: protocol_version.clone(), + }; + let mut server_peer_info = ServerInfo::default(); + server_peer_info.protocol_version = protocol_version; + + let server = serve_directly::( + CacheHintServer::new(), + server_transport, + Some(client_handler.get_info()), + ); let server_handle = tokio::spawn(async move { - CacheHintServer::new() - .serve(server_transport) - .await? - .waiting() - .await?; + server.waiting().await?; anyhow::Ok(()) }); - let client = VersionedClient { protocol_version } - .serve(client_transport) - .await - .expect("client should connect"); + let client = serve_directly::( + client_handler, + client_transport, + Some(server_peer_info.into()), + ); let tools = client .list_tools(None) .await diff --git a/crates/rmcp/tests/test_live_oauth_refresh.rs b/crates/rmcp/tests/test_live_oauth_refresh.rs new file mode 100644 index 000000000..543abe537 --- /dev/null +++ b/crates/rmcp/tests/test_live_oauth_refresh.rs @@ -0,0 +1,233 @@ +//! Live authorization-server checks for the refresh path. +//! +//! These are `#[ignore]`d, so they never run in CI. Provision the authorization +//! server with `./scripts/keycloak-oauth-fixture.sh`, then run them with +//! `cargo test -p rmcp --all-features --test test_live_oauth_refresh -- --ignored`. +//! `KC_BASE` overrides the server location for both the script and these tests. +#![cfg(feature = "auth")] + +use std::sync::Arc; + +use rmcp::transport::auth::{ + AuthError, AuthorizationManager, AuthorizationMetadata, CredentialRefreshGuard, + CredentialStore, InMemoryCredentialStore, OAuthClientConfig, OAuthTokenResponse, + StoredCredentials, +}; +use tokio::sync::Mutex; + +const REALM: &str = "rmcp"; +const CLIENT_ID: &str = "rmcp-client"; +const CLIENT_SECRET: &str = "rmcp-secret"; + +fn kc_base() -> String { + std::env::var("KC_BASE").unwrap_or_else(|_| "http://localhost:8081".to_string()) +} + +fn token_endpoint() -> String { + format!("{}/realms/{REALM}/protocol/openid-connect/token", kc_base()) +} + +/// Ask Keycloak for a genuine token pair through the direct access grant, so the +/// stored credentials hold a refresh token the server will actually honor. +async fn issue_real_credentials() -> OAuthTokenResponse { + let form = format!( + "client_id={CLIENT_ID}&client_secret={CLIENT_SECRET}\ + &username=alice&password=alice-pw&grant_type=password&scope=openid+profile" + ); + let body = reqwest::Client::new() + .post(token_endpoint()) + .header("content-type", "application/x-www-form-urlencoded") + .body(form) + .send() + .await + .expect("keycloak unreachable") + .text() + .await + .unwrap(); + serde_json::from_str(&body).unwrap_or_else(|e| panic!("unexpected token response {body}: {e}")) +} + +fn metadata() -> AuthorizationMetadata { + // `AuthorizationMetadata` is `#[non_exhaustive]`, so fill it field by field. + let mut metadata = AuthorizationMetadata::default(); + metadata.authorization_endpoint = + format!("{}/realms/{REALM}/protocol/openid-connect/auth", kc_base()); + metadata.token_endpoint = token_endpoint(); + metadata +} + +fn client_config() -> OAuthClientConfig { + let mut config = OAuthClientConfig::new(CLIENT_ID, "http://localhost/callback"); + config.client_secret = Some(CLIENT_SECRET.to_string()); + config +} + +async fn manager_with_store(store: S) -> AuthorizationManager { + let mut manager = AuthorizationManager::new(kc_base()).await.unwrap(); + manager.set_metadata(metadata()); + manager.configure_client(client_config()).unwrap(); + manager.set_credential_store(store); + manager +} + +fn stored(client_id: &str, token: OAuthTokenResponse) -> StoredCredentials { + StoredCredentials::new( + client_id.to_string(), + Some(token), + vec!["openid".into(), "profile".into()], + None, + ) +} + +/// A store that serializes refreshes the way a shared on-disk store would. +#[derive(Clone, Default)] +struct GuardedStore { + inner: InMemoryCredentialStore, + lock: Arc>, +} + +#[async_trait::async_trait] +impl CredentialStore for GuardedStore { + async fn load(&self) -> Result, AuthError> { + self.inner.load().await + } + + async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError> { + self.inner.save(credentials).await + } + + async fn clear(&self) -> Result<(), AuthError> { + self.inner.clear().await + } + + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + Ok(Some(CredentialRefreshGuard::new( + self.lock.clone().lock_owned().await, + ))) + } +} + +#[tokio::test] +#[ignore = "requires a live Keycloak"] +async fn live_refresh_rotates_the_stored_token() { + use oauth2::TokenResponse; + + let issued = issue_real_credentials().await; + let original_refresh = issued.refresh_token().unwrap().secret().clone(); + let store = InMemoryCredentialStore::new(); + store.save(stored(CLIENT_ID, issued)).await.unwrap(); + let manager = manager_with_store(store.clone()).await; + + let refreshed = manager.refresh_token().await.expect("live refresh failed"); + + let saved = store.load().await.unwrap().unwrap(); + let saved_token = saved.token_response.unwrap(); + assert_ne!( + saved_token.refresh_token().unwrap().secret(), + &original_refresh, + "keycloak rotates refresh tokens, so the store must hold the new one" + ); + assert_eq!( + saved_token.access_token().secret(), + refreshed.access_token().secret(), + "the saved credentials must match what the caller received" + ); + assert!( + saved.granted_scopes.contains(&"openid".to_string()), + "granted scopes should come from the provider response, got: {:?}", + saved.granted_scopes + ); +} + +#[tokio::test] +#[ignore = "requires a live Keycloak"] +async fn live_refresh_rejects_credentials_for_another_client() { + use oauth2::TokenResponse; + + let issued = issue_real_credentials().await; + let untouched_refresh = issued.refresh_token().unwrap().secret().clone(); + let store = InMemoryCredentialStore::new(); + store + .save(stored("some-other-client", issued)) + .await + .unwrap(); + let manager = manager_with_store(store.clone()).await; + + let error = manager.refresh_token().await.unwrap_err(); + + assert!( + matches!(error, AuthError::AuthorizationRequired), + "a client mismatch must require reauthorization, got: {error:?}" + ); + let saved = store.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + &untouched_refresh, + "a rejected refresh must leave the stored token untouched" + ); +} + +#[tokio::test] +#[ignore = "requires a live Keycloak"] +async fn live_concurrent_refreshes_survive_refresh_token_rotation() { + use oauth2::TokenResponse; + + let issued = issue_real_credentials().await; + let store = GuardedStore::default(); + store.save(stored(CLIENT_ID, issued)).await.unwrap(); + let first = manager_with_store(store.clone()).await; + let second = manager_with_store(store.clone()).await; + + // Without the guard the second caller would reuse the refresh token the first + // one already consumed, and Keycloak would answer `invalid_grant`. + let (a, b) = tokio::join!( + tokio::spawn(async move { first.refresh_token().await }), + tokio::spawn(async move { second.refresh_token().await }) + ); + + let a = a.unwrap().expect("first concurrent refresh failed"); + let b = b.unwrap().expect("second concurrent refresh failed"); + assert_ne!( + a.access_token().secret(), + b.access_token().secret(), + "each caller performs its own exchange, so the tokens must differ" + ); +} + +/// Deterministic proof that this realm really does invalidate a rotated refresh +/// token, which is what makes the coordination above load-bearing. +#[tokio::test] +#[ignore = "requires a live Keycloak with revokeRefreshToken enabled"] +async fn live_reusing_a_rotated_refresh_token_is_rejected() { + let issued = issue_real_credentials().await; + + let first_store = InMemoryCredentialStore::new(); + first_store + .save(stored(CLIENT_ID, issued.clone())) + .await + .unwrap(); + manager_with_store(first_store) + .await + .refresh_token() + .await + .expect("the first refresh should succeed"); + + // A second caller that never saw the rotation still holds the consumed token. + let stale_store = InMemoryCredentialStore::new(); + stale_store.save(stored(CLIENT_ID, issued)).await.unwrap(); + let error = manager_with_store(stale_store) + .await + .refresh_token() + .await + .unwrap_err(); + + assert!( + matches!(error, AuthError::TokenRefreshRejected(_)), + "reusing a rotated refresh token must be rejected, got: {error:?}" + ); +} diff --git a/crates/rmcp/tests/test_prompt_macros.rs b/crates/rmcp/tests/test_prompt_macros.rs index 7a00249a4..642ae88da 100644 --- a/crates/rmcp/tests/test_prompt_macros.rs +++ b/crates/rmcp/tests/test_prompt_macros.rs @@ -4,14 +4,12 @@ use std::sync::Arc; use rmcp::{ - ClientHandler, RoleServer, ServerHandler, ServiceExt, + ClientHandler, ServerHandler, ServiceExt, handler::server::{router::prompt::PromptRouter, wrapper::Parameters}, model::{ - ClientInfo, ContentBlock, GetPromptRequestParams, GetPromptResult, ListPromptsResult, - PaginatedRequestParams, PromptMessage, Role, + ClientInfo, ContentBlock, GetPromptRequestParams, GetPromptResult, PromptMessage, Role, }, prompt, prompt_handler, prompt_router, - service::RequestContext, }; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; diff --git a/crates/rmcp/tests/test_prompt_routers.rs b/crates/rmcp/tests/test_prompt_routers.rs index eecdb1e4a..e73057947 100644 --- a/crates/rmcp/tests/test_prompt_routers.rs +++ b/crates/rmcp/tests/test_prompt_routers.rs @@ -20,12 +20,6 @@ struct Request { fields: HashMap, } -#[derive(Debug, schemars::JsonSchema, serde::Deserialize, serde::Serialize)] -struct Sum { - a: i32, - b: i32, -} - #[rmcp::prompt_router(router = "test_router")] impl TestHandler { #[rmcp::prompt] diff --git a/crates/rmcp/tests/test_protocol_version_negotiation.rs b/crates/rmcp/tests/test_protocol_version_negotiation.rs index e91ecf97a..aa4607a1c 100644 --- a/crates/rmcp/tests/test_protocol_version_negotiation.rs +++ b/crates/rmcp/tests/test_protocol_version_negotiation.rs @@ -1,15 +1,25 @@ //! Tests for protocol version negotiation in the default ServerHandler::initialize impl. //! -//! Known versions are echoed back; unknown versions fall back to LATEST. +//! Handshake versions are echoed back; every other version falls back to one +//! the server can serve over `initialize`. #![cfg(not(feature = "local"))] #![cfg(feature = "client")] -use std::borrow::Cow; +use std::{ + borrow::Cow, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, +}; use rmcp::{ ClientHandler, ErrorData, RoleServer, ServerHandler, ServiceExt, - model::{ClientInfo, InitializeRequestParams, InitializeResult, ProtocolVersion, ServerInfo}, - service::RequestContext, + model::{ + ClientCapabilities, ClientInfo, ErrorCode, Implementation, InitializeRequestParams, + InitializeResult, ProtocolVersion, ServerInfo, + }, + service::{ClientInitializeError, RequestContext}, }; #[derive(Debug, Clone, Default)] @@ -21,9 +31,10 @@ impl ServerHandler for EchoServer { } } -/// Every known version except `2026-07-28`, standing in for a server that has -/// not implemented that revision. -const NARROWED_VERSIONS: &[ProtocolVersion] = &[ +/// Every known version whose lifecycle still runs the `initialize` handshake. +/// `2026-07-28` replaced the handshake with per-request metadata, so this is +/// also the list a server that has not implemented that revision supports. +const HANDSHAKE_VERSIONS: &[ProtocolVersion] = &[ ProtocolVersion::V_2024_11_05, ProtocolVersion::V_2025_03_26, ProtocolVersion::V_2025_06_18, @@ -39,7 +50,25 @@ impl ServerHandler for NarrowedServer { } fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { - Cow::Borrowed(NARROWED_VERSIONS) + Cow::Borrowed(HANDSHAKE_VERSIONS) + } +} + +/// Supports only revisions that have no `initialize` handshake at all. +#[derive(Debug, Clone, Default)] +struct ModernOnlyServer; + +const MODERN_ONLY_VERSIONS: &[ProtocolVersion] = &[ProtocolVersion::V_2026_07_28]; + +impl ServerHandler for ModernOnlyServer { + fn get_info(&self) -> ServerInfo { + let mut info = ServerInfo::default(); + info.protocol_version = ProtocolVersion::V_2026_07_28; + info + } + + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + Cow::Borrowed(MODERN_ONLY_VERSIONS) } } @@ -55,7 +84,7 @@ impl ServerHandler for NarrowedOverridingServer { } fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { - Cow::Borrowed(NARROWED_VERSIONS) + Cow::Borrowed(HANDSHAKE_VERSIONS) } async fn initialize( @@ -88,23 +117,28 @@ async fn negotiated_version_with( server: S, client_version: ProtocolVersion, ) -> ProtocolVersion { + negotiate_with(server, client_version) + .await + .expect("client should connect") +} + +async fn negotiate_with( + server: S, + client_version: ProtocolVersion, +) -> Result { let (server_transport, client_transport) = tokio::io::duplex(4096); tokio::spawn(async move { - let _ = server - .serve(server_transport) - .await - .expect("server should start") - .waiting() - .await; + if let Ok(running) = server.serve(server_transport).await { + let _ = running.waiting().await; + } }); let client = VersionedClient { protocol_version: client_version, } .serve(client_transport) - .await - .expect("client should connect"); + .await?; let version = client .peer_info() @@ -113,20 +147,55 @@ async fn negotiated_version_with( .clone(); client.cancel().await.expect("client should cancel"); - version + Ok(version) } #[tokio::test] -async fn known_version_echoed_back() { - for version in ProtocolVersion::KNOWN_VERSIONS { +async fn handshake_version_echoed_back() { + for version in HANDSHAKE_VERSIONS { let negotiated = negotiated_version(version.clone()).await; assert_eq!( negotiated, *version, - "known version {version} should be echoed back" + "handshake version {version} should be echoed back" ); } } +/// `initialize` disappeared in `2026-07-28`, so agreeing to it here would leave +/// the peers speaking a revision that has no handshake at all. +#[tokio::test] +async fn handshake_never_agrees_to_a_version_that_dropped_it() { + let negotiated = negotiated_version(ProtocolVersion::V_2026_07_28).await; + assert_eq!( + negotiated, + ProtocolVersion::LATEST, + "a version that dropped the handshake should fall back to the server's own" + ); +} + +#[tokio::test] +async fn modern_only_server_rejects_the_handshake() { + let error = negotiate_with(ModernOnlyServer, ProtocolVersion::V_2026_07_28) + .await + .expect_err("a server with no handshake version cannot answer initialize"); + let ClientInitializeError::JsonRpcError(error) = error else { + panic!("expected a JSON-RPC error, got {error:?}"); + }; + assert_eq!( + error.code, + ErrorCode::UNSUPPORTED_PROTOCOL_VERSION, + "a server with no handshake version should reject initialize" + ); + assert_eq!( + error.data, + Some(serde_json::json!({ + "requested": "2026-07-28", + "supported": ["2026-07-28"], + })), + "the rejection should name the versions the server does support" + ); +} + #[tokio::test] async fn unknown_version_falls_back_to_latest() { let unknown: ProtocolVersion = serde_json::from_str(r#""1999-01-01""#).unwrap(); @@ -140,7 +209,7 @@ async fn unknown_version_falls_back_to_latest() { #[tokio::test] async fn narrowed_server_still_echoes_versions_it_supports() { - for version in NARROWED_VERSIONS { + for version in HANDSHAKE_VERSIONS { let negotiated = negotiated_version_with(NarrowedServer, version.clone()).await; assert_eq!( negotiated, *version, @@ -169,3 +238,91 @@ async fn narrowed_server_caps_even_when_it_overrides_initialize() { "the handshake layer should not raise the version above what the server supports" ); } + +/// Overrides `initialize` to run a side effect, then delegates the version +/// answer back to the SDK with [`ServerHandler::negotiate_initialize`]. +#[derive(Debug, Clone, Default)] +struct DelegatingServer { + initializations: Arc, +} + +impl ServerHandler for DelegatingServer { + fn get_info(&self) -> ServerInfo { + ServerInfo::default() + } + + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + Cow::Borrowed(HANDSHAKE_VERSIONS) + } + + async fn initialize( + &self, + request: InitializeRequestParams, + context: RequestContext, + ) -> Result { + self.initializations.fetch_add(1, Ordering::Relaxed); + context.peer.set_peer_info(request.clone()); + self.negotiate_initialize(&request) + } +} + +fn initialize_params(protocol_version: ProtocolVersion) -> InitializeRequestParams { + let mut params = InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("test-client", "0.0.0"), + ); + params.protocol_version = protocol_version; + params +} + +#[test] +fn negotiate_initialize_echoes_a_supported_version() { + let result = NarrowedServer + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2025_06_18)) + .expect("a supported handshake version should negotiate"); + assert_eq!(result.protocol_version, ProtocolVersion::V_2025_06_18); +} + +#[test] +fn negotiate_initialize_caps_at_supported_versions() { + let result = NarrowedServer + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) + .expect("an unsupported version should fall back rather than fail"); + assert_eq!(result.protocol_version, ProtocolVersion::V_2025_11_25); +} + +#[test] +fn negotiate_initialize_keeps_the_rest_of_get_info() { + let server = NarrowedServer; + let result = server + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) + .expect("an unsupported version should fall back rather than fail"); + assert_eq!(result.capabilities, server.get_info().capabilities); +} + +#[test] +fn negotiate_initialize_rejects_when_no_handshake_version_is_supported() { + let error = ModernOnlyServer + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) + .expect_err("a server with no handshake version cannot answer initialize"); + assert_eq!(error.code, ErrorCode::UNSUPPORTED_PROTOCOL_VERSION); +} + +#[tokio::test] +async fn delegating_server_negotiates_like_the_default_initialize() { + let negotiated = + negotiated_version_with(DelegatingServer::default(), ProtocolVersion::V_2026_07_28).await; + assert_eq!( + negotiated, + ProtocolVersion::V_2025_11_25, + "an override that delegates should answer what the default initialize would" + ); +} + +#[tokio::test] +async fn delegating_server_still_runs_its_own_side_effect() { + let server = DelegatingServer::default(); + let initializations = Arc::clone(&server.initializations); + negotiated_version_with(server, ProtocolVersion::V_2025_06_18).await; + assert_eq!(initializations.load(Ordering::Relaxed), 1); +} diff --git a/crates/rmcp/tests/test_resource_not_found_version.rs b/crates/rmcp/tests/test_resource_not_found_version.rs index 44eb3631f..1d86ca8c3 100644 --- a/crates/rmcp/tests/test_resource_not_found_version.rs +++ b/crates/rmcp/tests/test_resource_not_found_version.rs @@ -6,12 +6,12 @@ #![cfg(feature = "client")] use rmcp::{ - ClientHandler, RoleServer, ServerHandler, ServiceError, ServiceExt, + ClientHandler, RoleClient, RoleServer, ServerHandler, ServiceError, model::{ ClientInfo, ErrorCode, ErrorData, ProtocolVersion, ReadResourceRequestParams, - ReadResourceResponse, + ReadResourceResponse, ServerInfo, }, - service::RequestContext, + service::{RequestContext, serve_directly}, }; #[derive(Debug, Clone, Default)] @@ -40,24 +40,33 @@ impl ClientHandler for VersionedClient { } } +/// Wires the pair up directly on `client_version`. `2026-07-28` removed the +/// `initialize` handshake, so a peer on that revision is reached the way the +/// discover lifecycle leaves one: with the version already agreed. async fn not_found_code(client_version: ProtocolVersion) -> ErrorCode { let (server_transport, client_transport) = tokio::io::duplex(4096); + let client_handler = VersionedClient { + protocol_version: client_version.clone(), + }; + let mut server_peer_info = ServerInfo::default(); + server_peer_info.protocol_version = client_version; + + let server = serve_directly::( + ResourceServer, + server_transport, + Some(client_handler.get_info()), + ); let server_handle = tokio::spawn(async move { - ResourceServer - .serve(server_transport) - .await? - .waiting() - .await?; + server.waiting().await?; anyhow::Ok(()) }); - let client = VersionedClient { - protocol_version: client_version, - } - .serve(client_transport) - .await - .expect("client should connect"); + let client = serve_directly::( + client_handler, + client_transport, + Some(server_peer_info.into()), + ); let error = client .read_resource(ReadResourceRequestParams::new("missing://resource")) diff --git a/crates/rmcp/tests/test_result_type_version.rs b/crates/rmcp/tests/test_result_type_version.rs index 849a6610e..ab779d283 100644 --- a/crates/rmcp/tests/test_result_type_version.rs +++ b/crates/rmcp/tests/test_result_type_version.rs @@ -6,12 +6,12 @@ #![cfg(feature = "client")] use rmcp::{ - ClientHandler, RoleServer, ServerHandler, ServiceExt, + ClientHandler, RoleClient, RoleServer, ServerHandler, model::{ CallToolRequestParams, CallToolResponse, CallToolResult, ClientInfo, ContentBlock, - ErrorData, ProtocolVersion, ResultType, + ErrorData, ProtocolVersion, ResultType, ServerInfo, }, - service::RequestContext, + service::{RequestContext, serve_directly}, }; #[derive(Debug, Clone, Default)] @@ -40,20 +40,33 @@ impl ClientHandler for VersionedClient { } } +/// Wires the pair up directly on `client_version`. `2026-07-28` removed the +/// `initialize` handshake, so a peer on that revision is reached the way the +/// discover lifecycle leaves one: with the version already agreed. async fn call_tool_result_type(client_version: ProtocolVersion) -> Option { let (server_transport, client_transport) = tokio::io::duplex(4096); + let client_handler = VersionedClient { + protocol_version: client_version.clone(), + }; + let mut server_peer_info = ServerInfo::default(); + server_peer_info.protocol_version = client_version; + + let server = serve_directly::( + ToolServer, + server_transport, + Some(client_handler.get_info()), + ); let server_handle = tokio::spawn(async move { - ToolServer.serve(server_transport).await?.waiting().await?; + server.waiting().await?; anyhow::Ok(()) }); - let client = VersionedClient { - protocol_version: client_version, - } - .serve(client_transport) - .await - .expect("client should connect"); + let client = serve_directly::( + client_handler, + client_transport, + Some(server_peer_info.into()), + ); let result = client .call_tool(CallToolRequestParams::new("echo")) diff --git a/crates/rmcp/tests/test_sampling.rs b/crates/rmcp/tests/test_sampling.rs index b108ff412..3c82fd42b 100644 --- a/crates/rmcp/tests/test_sampling.rs +++ b/crates/rmcp/tests/test_sampling.rs @@ -375,7 +375,7 @@ fn test_tool_result_content_requires_content() { #[case::array(serde_json::json!([{ "city": "SF", "temp": 72 }, { "city": "NY", "temp": 65 }]))] #[case::string(serde_json::json!("sunny"))] #[case::integer(serde_json::json!(42))] -#[case::float(serde_json::json!(3.14))] +#[case::float(serde_json::json!(3.5))] #[case::boolean(serde_json::json!(true))] fn tool_result_content_round_trips_non_object_structured_content( #[case] structured: serde_json::Value, diff --git a/crates/rmcp/tests/test_sep_2260_request_association.rs b/crates/rmcp/tests/test_sep_2260_request_association.rs index d4e20e1e9..b02cff037 100644 --- a/crates/rmcp/tests/test_sep_2260_request_association.rs +++ b/crates/rmcp/tests/test_sep_2260_request_association.rs @@ -1,5 +1,8 @@ #![cfg(all(feature = "server", feature = "client", not(feature = "local")))] -#![allow(deprecated)] +#![expect( + deprecated, + reason = "This test verifies request association for the deprecated sampling API" +)] use std::sync::{Arc, Mutex}; @@ -10,7 +13,7 @@ use rmcp::{ CreateMessageRequest, CreateMessageRequestParams, CreateMessageResult, ProtocolVersion, SamplingMessage, ServerCapabilities, ServerInfo, ServerRequest, }, - service::RequestContext, + service::{RequestContext, RunningService, serve_directly}, }; use serde_json::{Value, json}; use tokio::{ @@ -18,9 +21,11 @@ use tokio::{ sync::oneshot, }; +type RequestResultSender = oneshot::Sender>; + #[derive(Clone)] struct SamplingServer { - outside: Arc>>>>, + outside: Arc>>, } impl ServerHandler for SamplingServer { @@ -95,21 +100,44 @@ impl ClientHandler for SamplingClient { } } +/// Connects the pair on `2026-07-28`. That revision dropped the `initialize` +/// handshake, so the version is agreed up front the way a discover-lifecycle +/// startup leaves it. +fn serve_modern_pair( + server: SamplingServer, +) -> ( + RunningService, + RunningService, +) { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let mut server_peer_info = server.get_info(); + server_peer_info.protocol_version = ProtocolVersion::V_2026_07_28; + + let running_server = serve_directly::( + server, + server_transport, + Some(SamplingClient.get_info()), + ); + let client = serve_directly::( + SamplingClient, + client_transport, + Some(server_peer_info.into()), + ); + (running_server, client) +} + #[tokio::test] async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> { - let (server_transport, client_transport) = tokio::io::duplex(4096); let (tx, rx) = oneshot::channel(); let server = SamplingServer { outside: Arc::new(Mutex::new(Some(tx))), }; + let (running_server, client) = serve_modern_pair(server); let server_handle = tokio::spawn(async move { - let running = server.serve(server_transport).await?; - running.waiting().await?; + running_server.waiting().await?; anyhow::Ok(()) }); - let client = SamplingClient.serve(client_transport).await?; - let result = client .peer() .call_tool(CallToolRequestParams::new("sample")) @@ -129,19 +157,16 @@ async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> { #[tokio::test] async fn generic_send_request_bypass_rejected() -> anyhow::Result<()> { - let (server_transport, client_transport) = tokio::io::duplex(4096); let (tx, rx) = oneshot::channel(); let server = SamplingServer { outside: Arc::new(Mutex::new(Some(tx))), }; + let (running_server, client) = serve_modern_pair(server); let server_handle = tokio::spawn(async move { - let running = server.serve(server_transport).await?; - running.waiting().await?; + running_server.waiting().await?; anyhow::Ok(()) }); - let client = SamplingClient.serve(client_transport).await?; - let result = client .peer() .call_tool(CallToolRequestParams::new("sample_generic")) diff --git a/crates/rmcp/tests/test_stateless_protocol_version.rs b/crates/rmcp/tests/test_stateless_protocol_version.rs index 02222ec2c..dbfcaf27c 100644 --- a/crates/rmcp/tests/test_stateless_protocol_version.rs +++ b/crates/rmcp/tests/test_stateless_protocol_version.rs @@ -1,7 +1,8 @@ //! Tests for protocol version negotiation in stateless HTTP mode. //! -//! Supported versions are echoed back; unknown versions, and versions outside -//! the server's `supported_protocol_versions`, fall back to the handler's own +//! Supported handshake versions are echoed back; unknown versions, versions +//! outside the server's `supported_protocol_versions`, and versions that no +//! longer have an `initialize` handshake fall back to the handler's own //! version. #![cfg(not(feature = "local"))] @@ -36,9 +37,10 @@ impl ServerHandler for OverridingInitialize { } } -/// Every known version except `2026-07-28`, standing in for a server that has -/// not implemented that revision. -const NARROWED_VERSIONS: &[ProtocolVersion] = &[ +/// Every known version whose lifecycle still runs the `initialize` handshake. +/// `2026-07-28` replaced the handshake with per-request metadata, so this is +/// also the list a server that has not implemented that revision supports. +const HANDSHAKE_VERSIONS: &[ProtocolVersion] = &[ ProtocolVersion::V_2024_11_05, ProtocolVersion::V_2025_03_26, ProtocolVersion::V_2025_06_18, @@ -56,7 +58,7 @@ impl ServerHandler for NarrowedOverridingInitialize { } fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { - Cow::Borrowed(NARROWED_VERSIONS) + Cow::Borrowed(HANDSHAKE_VERSIONS) } async fn initialize( @@ -148,15 +150,15 @@ async fn post_init(client: &reqwest::Client, url: &str, body_version: &str) -> s } #[tokio::test] -async fn stateless_json_init_echoes_known_versions_when_handler_overrides_initialize() { +async fn stateless_json_init_echoes_handshake_versions_when_handler_overrides_initialize() { let (client, url, ct) = spawn_server(stateless_json_config()).await; - for version in ProtocolVersion::KNOWN_VERSIONS { + for version in HANDSHAKE_VERSIONS { let resp = post_init(&client, &url, version.as_str()).await; assert_eq!( resp["result"]["protocolVersion"], version.as_str(), - "known version {version} should be echoed back" + "handshake version {version} should be echoed back" ); } @@ -164,15 +166,15 @@ async fn stateless_json_init_echoes_known_versions_when_handler_overrides_initia } #[tokio::test] -async fn stateless_sse_init_echoes_known_versions_when_handler_overrides_initialize() { +async fn stateless_sse_init_echoes_handshake_versions_when_handler_overrides_initialize() { let (client, url, ct) = spawn_server(stateless_sse_config()).await; - for version in ProtocolVersion::KNOWN_VERSIONS { + for version in HANDSHAKE_VERSIONS { let resp = post_init(&client, &url, version.as_str()).await; assert_eq!( resp["result"]["protocolVersion"], version.as_str(), - "known version {version} should be echoed back" + "handshake version {version} should be echoed back" ); } @@ -198,7 +200,7 @@ async fn stateless_json_init_echoes_versions_the_server_narrowed_to() { let (client, url, ct) = spawn_server_of::(stateless_json_config()).await; - for version in NARROWED_VERSIONS { + for version in HANDSHAKE_VERSIONS { let resp = post_init(&client, &url, version.as_str()).await; assert_eq!( resp["result"]["protocolVersion"], diff --git a/crates/rmcp/tests/test_streamable_http_protocol_version.rs b/crates/rmcp/tests/test_streamable_http_protocol_version.rs index fcbeb41a0..373549aa9 100644 --- a/crates/rmcp/tests/test_streamable_http_protocol_version.rs +++ b/crates/rmcp/tests/test_streamable_http_protocol_version.rs @@ -87,6 +87,17 @@ async fn post_init( req.send().await.expect("send initialize request") } +/// First JSON-RPC message of an SSE response, skipping the empty priming event. +async fn sse_payload(response: reqwest::Response) -> anyhow::Result { + let body = response.text().await?; + let data = body + .lines() + .filter_map(|line| line.strip_prefix("data: ")) + .find(|data| !data.is_empty()) + .expect("response must carry a JSON-RPC payload"); + Ok(serde_json::from_str(data)?) +} + async fn post_non_initialize(client: &reqwest::Client, url: &str) -> reqwest::Response { client .post(url) @@ -186,6 +197,27 @@ async fn stateless_init_accepts_when_header_absent() -> anyhow::Result<()> { Ok(()) } +#[tokio::test] +async fn stateless_init_naming_a_post_handshake_version_negotiates_down() -> anyhow::Result<()> { + let (client, url, ct) = spawn_server(stateless_json_config()).await; + + let response = post_init(&client, &url, None, "2026-07-28").await; + assert_eq!( + response.status(), + 200, + "initialize selects legacy semantics whatever version it names" + ); + + let body: serde_json::Value = response.json().await?; + assert_eq!( + body["result"]["protocolVersion"], "2025-11-25", + "initialize must settle on a version that still has the handshake" + ); + + ct.cancel(); + Ok(()) +} + #[tokio::test] async fn stateful_init_rejects_when_header_mismatches_body() -> anyhow::Result<()> { let (client, url, ct) = spawn_server(stateful_config()).await; @@ -218,6 +250,37 @@ async fn stateful_rejected_initial_posts_do_not_create_sessions() -> anyhow::Res Ok(()) } +/// The `initialize` handshake exists only in the revisions before `2026-07-28`, +/// so naming a later version in it does not make the request a modern one: the +/// server keeps legacy semantics, opens a session, and answers with a version it +/// can actually serve over the handshake. +#[tokio::test] +async fn stateful_init_naming_a_post_handshake_version_opens_a_legacy_session() -> anyhow::Result<()> +{ + let (client, url, ct) = spawn_server(stateful_config()).await; + + let response = post_init(&client, &url, None, "2026-07-28").await; + assert_eq!( + response.status(), + 200, + "initialize selects legacy semantics whatever version it names" + ); + assert!( + response.headers().contains_key("Mcp-Session-Id"), + "the handshake must open a session, got headers: {:?}", + response.headers() + ); + + let payload = sse_payload(response).await?; + assert_eq!( + payload["result"]["protocolVersion"], "2025-11-25", + "initialize must settle on a version that still has the handshake" + ); + + ct.cancel(); + Ok(()) +} + #[tokio::test] async fn stateless_missing_protocol_header_returns_header_mismatch() -> anyhow::Result<()> { let (client, url, ct) = spawn_server(stateless_json_config()).await; diff --git a/crates/rmcp/tests/test_task.rs b/crates/rmcp/tests/test_task.rs index ea1a2595e..55ba4ccb2 100644 --- a/crates/rmcp/tests/test_task.rs +++ b/crates/rmcp/tests/test_task.rs @@ -13,9 +13,9 @@ use rmcp::{ use serde_json::json; #[derive(Debug, serde::Deserialize, rmcp::schemars::JsonSchema)] -pub struct SumArgs { - pub a: i32, - pub b: i32, +struct SumArgs { + a: i32, + b: i32, } #[derive(Clone)] diff --git a/crates/rmcp/tests/test_tool_macros.rs b/crates/rmcp/tests/test_tool_macros.rs index 9b9530aa8..4975109cb 100644 --- a/crates/rmcp/tests/test_tool_macros.rs +++ b/crates/rmcp/tests/test_tool_macros.rs @@ -573,3 +573,47 @@ fn test_manual_get_info_not_overridden() { "manual resources should be preserved" ); } + +/// Server whose tools come from a `macro_rules!` helper wrapping the whole annotated impl. +#[derive(Debug, Clone)] +struct MacroGeneratedServer; + +macro_rules! define_tools { + ($($name:ident => $description:literal),* $(,)?) => { + #[tool_router] + impl MacroGeneratedServer { + $( + #[tool(description = $description)] + async fn $name(&self) -> String { + stringify!($name).to_owned() + } + )* + } + }; +} + +define_tools!(probe => "what a capability would own"); + +#[test] +fn test_macro_rules_around_the_impl_registers_tools() { + let tools = MacroGeneratedServer::tool_router().list_all(); + + assert_eq!(tools.len(), 1); + assert_eq!(tools[0].name, "probe"); + assert_eq!( + tools[0].description.as_deref(), + Some("what a capability would own") + ); +} + +/// Server that opts in to a router with no tools. +#[derive(Debug, Clone)] +struct EmptyRouterServer; + +#[tool_router(allow_empty)] +impl EmptyRouterServer {} + +#[test] +fn test_allow_empty_builds_a_router_without_tools() { + assert!(EmptyRouterServer::tool_router().list_all().is_empty()); +} diff --git a/crates/rmcp/tests/test_tool_routers.rs b/crates/rmcp/tests/test_tool_routers.rs index d2bbe8687..12a5d4000 100644 --- a/crates/rmcp/tests/test_tool_routers.rs +++ b/crates/rmcp/tests/test_tool_routers.rs @@ -21,12 +21,6 @@ struct Request { fields: HashMap, } -#[derive(Debug, schemars::JsonSchema, serde::Deserialize, serde::Serialize)] -struct Sum { - a: i32, - b: i32, -} - #[rmcp::tool_router(router = test_router_1)] impl TestHandler { #[rmcp::tool] diff --git a/crates/rmcp/tests/test_typed_child_process.rs b/crates/rmcp/tests/test_typed_child_process.rs new file mode 100644 index 000000000..9ff26b816 --- /dev/null +++ b/crates/rmcp/tests/test_typed_child_process.rs @@ -0,0 +1,54 @@ +#![cfg(all( + unix, + not(feature = "local"), + feature = "transport-child-process", + feature = "client" +))] + +use rmcp::{ + ServiceExt, + model::{ClientRequest, CustomRequest, PingRequest, ServerResult}, + transport::{ConfigureCommandExt, TokioChildProcess}, +}; +use serde::Deserialize; + +#[derive(Debug, Deserialize, PartialEq)] +#[serde(rename_all = "camelCase")] +struct SkillsListResult { + result_type: String, + skills: Vec, + #[serde(rename = "_meta")] + meta: serde_json::Value, +} + +#[tokio::test] +async fn typed_extension_response_survives_child_process_and_preserves_correlation() +-> anyhow::Result<()> { + let fixture = concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/support/typed_stdio_server.sh" + ); + let transport = + TokioChildProcess::new(tokio::process::Command::new("sh").configure(|command| { + command.arg(fixture); + }))?; + let client = ().serve(transport).await?; + + let response: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await?; + assert_eq!(response.result_type, "complete"); + assert_eq!(response.skills, ["stdio-example"]); + assert_eq!(response.meta["vendorExtension"]["retained"], true); + + let following = client + .send_request(ClientRequest::PingRequest(PingRequest::default())) + .await?; + assert!(matches!(following, ServerResult::EmptyResult(_))); + + client.cancel().await?; + Ok(()) +} diff --git a/crates/rmcp/tests/test_unix_socket_transport.rs b/crates/rmcp/tests/test_unix_socket_transport.rs index 4c4ad52f1..eaaf2d858 100644 --- a/crates/rmcp/tests/test_unix_socket_transport.rs +++ b/crates/rmcp/tests/test_unix_socket_transport.rs @@ -13,11 +13,13 @@ use http::{HeaderName, HeaderValue}; use hyper_util::rt::TokioIo; use rmcp::{ ServiceExt, + model::{ClientRequest, CustomRequest}, transport::{ StreamableHttpClientTransport, UnixSocketHttpClient, streamable_http_client::StreamableHttpClientTransportConfig, }, }; +use serde::Deserialize; use serde_json::json; use tokio::sync::Mutex; @@ -35,10 +37,10 @@ async fn mcp_handler( let mut headers_map = HashMap::new(); for (name, value) in headers.iter() { let name_str = name.as_str(); - if name_str.starts_with("x-") || name_str == "host" { - if let Ok(v) = value.to_str() { - headers_map.insert(name_str.to_string(), v.to_string()); - } + if (name_str.starts_with("x-") || name_str == "host") + && let Ok(v) = value.to_str() + { + headers_map.insert(name_str.to_string(), v.to_string()); } } @@ -46,46 +48,67 @@ async fn mcp_handler( stored.extend(headers_map); drop(stored); - if let Ok(json_body) = serde_json::from_slice::(&body) { - if let Some(method) = json_body.get("method").and_then(|m| m.as_str()) { - if method == "initialize" { - state.initialize_called.notify_one(); - let response = json!({ - "jsonrpc": "2.0", - "id": json_body.get("id"), - "result": { - "protocolVersion": "2024-11-05", - "capabilities": {}, - "serverInfo": { - "name": "test-unix-server", - "version": "1.0.0" - } + if let Ok(json_body) = serde_json::from_slice::(&body) + && let Some(method) = json_body.get("method").and_then(|m| m.as_str()) + { + if method == "initialize" { + state.initialize_called.notify_one(); + let response = json!({ + "jsonrpc": "2.0", + "id": json_body.get("id"), + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "serverInfo": { + "name": "test-unix-server", + "version": "1.0.0" } - }); - return ( - StatusCode::OK, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "unix-test-session", - ), - ], - response.to_string(), - ); - } else if method == "notifications/initialized" { - return ( - StatusCode::ACCEPTED, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "unix-test-session", - ), - ], - String::new(), - ); - } + } + }); + return ( + StatusCode::OK, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "unix-test-session", + ), + ], + response.to_string(), + ); + } else if method == "notifications/initialized" { + return ( + StatusCode::ACCEPTED, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "unix-test-session", + ), + ], + String::new(), + ); + } else if method == "skills/list" { + let response = json!({ + "jsonrpc": "2.0", + "id": json_body.get("id"), + "result": { + "resultType": "complete", + "skills": ["unix-example"], + "_meta": {"vendorExtension": {"retained": true}} + } + }); + return ( + StatusCode::OK, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "unix-test-session", + ), + ], + response.to_string(), + ); } } @@ -111,6 +134,81 @@ async fn mcp_handler( ) } +#[derive(Debug, Deserialize, PartialEq)] +#[serde(rename_all = "camelCase")] +struct SkillsListResult { + result_type: String, + skills: Vec, + #[serde(rename = "_meta")] + meta: serde_json::Value, +} + +struct TemporarySocketDirectory(std::path::PathBuf); + +impl TemporarySocketDirectory { + fn new() -> std::io::Result { + // Keep the pathname below the small sockaddr_un limit on macOS; its + // resolved system temporary directory can itself be very long. + let path = std::path::Path::new("/tmp").join(format!("rmcp-{}", uuid::Uuid::new_v4())); + std::fs::create_dir(&path)?; + Ok(Self(path)) + } +} + +impl Drop for TemporarySocketDirectory { + fn drop(&mut self) { + let _ = std::fs::remove_dir_all(&self.0); + } +} + +struct AbortServerOnDrop(tokio::task::JoinHandle<()>); + +impl Drop for AbortServerOnDrop { + fn drop(&mut self) { + self.0.abort(); + } +} + +/// Typed extension results must retain fields that the core response union does +/// not know about when transported over a Unix-domain HTTP connection. +#[tokio::test] +async fn test_unix_socket_typed_custom_response_preserves_extension_fields() -> anyhow::Result<()> { + let dir = TemporarySocketDirectory::new()?; + let socket_path = dir.0.join("mcp.sock"); + + let state = ServerState { + received_headers: Arc::new(Mutex::new(HashMap::new())), + initialize_called: Arc::new(tokio::sync::Notify::new()), + }; + let app = Router::new() + .route("/mcp", post(mcp_handler)) + .with_state(state); + let listener = tokio::net::UnixListener::bind(&socket_path)?; + let _server_guard = AbortServerOnDrop(spawn_unix_server(listener, app)); + + let socket_str = socket_path.to_str().expect("UTF-8 temporary path"); + let uri = "http://mcp-server.internal/mcp"; + let transport = StreamableHttpClientTransport::with_client( + UnixSocketHttpClient::new(socket_str, uri), + StreamableHttpClientTransportConfig::with_uri(uri), + ); + let client = ().serve(transport).await?; + + let response: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await?; + + assert_eq!(response.result_type, "complete"); + assert_eq!(response.skills, ["unix-example"]); + assert_eq!(response.meta["vendorExtension"]["retained"], true); + + client.cancel().await?; + Ok(()) +} + /// Spawns an HTTP/1.1 server on a Unix socket using hyper directly. /// Avoids `axum::serve(UnixListener, ...)` which uses `spawn_local` on Linux. fn spawn_unix_server( diff --git a/docs/OAUTH_SUPPORT.md b/docs/OAUTH_SUPPORT.md index 82a1fe918..5d32be348 100644 --- a/docs/OAUTH_SUPPORT.md +++ b/docs/OAUTH_SUPPORT.md @@ -15,6 +15,7 @@ This document describes the OAuth 2.1 authorization implementation for Model Con - Automatic token refresh - Authorized HTTP Client implementation - Injectable OAuth HTTP client for custom network environments +- Opt-in EMA/XAA refresh-token and ID-JAG exchanges for registered public and confidential clients ## Usage Guide @@ -294,6 +295,130 @@ match oauth_state.request_scope_upgrade("admin:write", MCP_REDIRECT_URI).await { } ``` +## Enterprise-managed authorization (EMA/XAA) + +The example requires the `rmcp` features `auth-enterprise-managed`, `client`, +`reqwest` (TLS), and `transport-streamable-http-client-reqwest`, plus `oauth2` +version 5. Call the async function from a Tokio runtime. + +The exchange profile has these requirements and limits: + +- Each authorization server has its own approved client registration and explicit + `EmaClientAuthentication`: `None`, `ClientSecretBasic`, `ClientSecretPost`, or + `JwtAssertion`. The SDK does not select methods from metadata or fall back to a + different method after a failure. +- Input is an enterprise IdP refresh token. The requested MCP resource must match + the ID-JAG's sole `resource` claim; scope may be omitted or narrowed. +- RAR (`authorization_details`) and DPoP are not supported. Nonempty authorization + details are rejected at both exchange stages. +- Redemption consumes the SDK's ID-JAG handle and does not retry automatically. + This is an SDK safety choice, not a protocol requirement that ID-JAGs be single-use. + +The [ID-JAG draft recommends confidential clients](https://datatracker.ietf.org/doc/html/draft-ietf-oauth-identity-assertion-authz-grant-04#section-9.1). +The example below uses confidential clients with `client_secret_basic` at both +servers. Use `None` only where that server permits a public client registration. +Discovery, server approval, SSO, credential storage, and reauthentication remain +the application's responsibility. Client-side ID-JAG checks validate structure and +bindings, not signatures; the resource authorization server verifies signatures. + +```rust no_run +use oauth2::{ClientSecret, RefreshToken}; +use rmcp::{ + ServiceExt, + model::ClientInfo, + transport::{ + StreamableHttpClientTransport, + auth::{ + default_oauth_http_client, + enterprise::{EmaAuthorizationServer, EmaClientAuthentication, EmaExchangeRequest}, + }, + streamable_http_client::StreamableHttpClientTransportConfig, + }, +}; + +async fn connect( + refresh: &RefreshToken, + idp_client_secret: ClientSecret, + resource_client_secret: ClientSecret, +) -> Result<(), Box> { + let resource = "https://mcp.example/mcp"; + let http = default_oauth_http_client()?; + let idp = EmaAuthorizationServer::new( + "https://idp.example", "https://idp.example/token", "idp-client", + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(idp_client_secret)); + let resource_as = EmaAuthorizationServer::new( + "https://as.example", "https://as.example/token", "mcp-client", + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(resource_client_secret)); + let token = EmaExchangeRequest::new(idp, resource_as, resource, refresh) + .with_scopes(["files.read"]) + .exchange(&http, &http) + .await?; + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(resource) + .auth_header(token.access_token.secret()), + ); + let client = ClientInfo::default().serve(transport).await?; + client.list_tools(Default::default()).await?; + client.cancel().await?; + Ok(()) +} +``` + +`auth_header` takes the token without a `Bearer ` prefix. Use it only with the +approved resource and never log it. This transport uses a fixed token; obtain a +new token and reconnect when it expires or is rejected. + +For a registration using JWT client authentication, implement +`EmaClientAssertionProvider` with your application's signer and configure +`JwtAssertion` on that server. The provider example below also uses `async-trait` +version 0.1; `AppSigner` represents your application's existing signing service. + +```rust ignore +use std::sync::Arc; +use rmcp::transport::auth::enterprise::{ + EmaAuthorizationServer, EmaClientAssertion, EmaClientAssertionProvider, + EmaClientAuthentication, +}; + +struct AppAssertionProvider { + signer: AppSigner, +} + +#[async_trait::async_trait] +impl EmaClientAssertionProvider for AppAssertionProvider { + async fn create_assertion( + &self, + server: &EmaAuthorizationServer, + ) -> Result> { + // Sign a new assertion for this registration and approved server. + let jwt = self.signer.sign_client_assertion( + &server.client_id, &server.issuer, &server.token_endpoint, + ).await?; + Ok(EmaClientAssertion::new(jwt)) + } +} + +let resource_as = EmaAuthorizationServer::new( + "https://as.example", "https://as.example/token", "mcp-client", +) +.with_client_authentication(EmaClientAuthentication::JwtAssertion(Arc::new( + AppAssertionProvider { signer }, +))); +``` + +The SDK calls the provider before each token request, including delayed +`EmaIdJag::exchange` redemption. Sign a fresh, short-lived assertion with a unique +`jti`, the registered client ID in `iss` and `sub`, and the server's approved +audience in `aud`. Your signer owns the keys and algorithm; the client assertion +is separate from the ID-JAG grant. Signing and HTTP share a 30-second deadline. + +The factory honors per-request redirect policy with the SDK's default reqwest +settings. For custom proxy, CA, or remote-execution policy, implement +`OAuthHttpClient`; use separate adapters for the IdP and resource AS when their +network policies differ. + ## Complete Examples - **Authorization Code client**: [`examples/clients/src/auth/oauth_client.rs`](../examples/clients/src/auth/oauth_client.rs) diff --git a/examples/clients/src/progress_client.rs b/examples/clients/src/progress_client.rs index a9f68c139..ad00dea32 100644 --- a/examples/clients/src/progress_client.rs +++ b/examples/clients/src/progress_client.rs @@ -100,10 +100,10 @@ impl ProgressAwareClient { } fn stop_tracking(&self) { - if let Ok(mut tracker_opt) = self.tracker.lock() { - if let Some(tracker) = tracker_opt.take() { - tracker.print_summary(); - } + if let Ok(mut tracker_opt) = self.tracker.lock() + && let Some(tracker) = tracker_opt.take() + { + tracker.print_summary(); } } } @@ -114,10 +114,10 @@ impl ClientHandler for ProgressAwareClient { params: ProgressNotificationParam, _context: NotificationContext, ) { - if let Ok(tracker_opt) = self.tracker.lock() { - if let Some(tracker) = tracker_opt.as_ref() { - tracker.handle_progress(¶ms); - } + if let Ok(tracker_opt) = self.tracker.lock() + && let Some(tracker) = tracker_opt.as_ref() + { + tracker.handle_progress(¶ms); } } @@ -182,10 +182,10 @@ async fn test_stdio_transport(records: u32) -> Result<()> { .call_tool(CallToolRequestParams::new("stream_processor")) .await?; - if let Some(content) = tool_result.content.first() { - if let Some(text) = content.as_text() { - tracing::info!("Processing completed: {}", text.text); - } + if let Some(content) = tool_result.content.first() + && let Some(text) = content.as_text() + { + tracing::info!("Processing completed: {}", text.text); } service.cancel().await?; @@ -236,10 +236,10 @@ async fn test_http_transport(http_url: &str, records: u32) -> Result<()> { .call_tool(CallToolRequestParams::new("stream_processor")) .await?; - if let Some(content) = tool_result.content.first() { - if let Some(text) = content.as_text() { - tracing::info!("processing completed: {}", text.text); - } + if let Some(content) = tool_result.content.first() + && let Some(text) = content.as_text() + { + tracing::info!("processing completed: {}", text.text); } client.cancel().await?; diff --git a/examples/clients/src/sampling_stdio.rs b/examples/clients/src/sampling_stdio.rs index cc7c5f153..107958624 100644 --- a/examples/clients/src/sampling_stdio.rs +++ b/examples/clients/src/sampling_stdio.rs @@ -1,3 +1,8 @@ +#![expect( + deprecated, + reason = "This example demonstrates the deprecated MCP sampling API" +)] + use anyhow::Result; use rmcp::{ ClientHandler, ServiceExt, diff --git a/examples/servers/src/cimd_auth_streamhttp.rs b/examples/servers/src/cimd_auth_streamhttp.rs index 6a634d883..a81297715 100644 --- a/examples/servers/src/cimd_auth_streamhttp.rs +++ b/examples/servers/src/cimd_auth_streamhttp.rs @@ -151,16 +151,16 @@ async fn fetch_and_validate_client_metadata(client_id_url: &str) -> Result, ) -> Result { if let Some(http_request_part) = context.extensions.get::() { @@ -278,7 +278,8 @@ impl ServerHandler for Counter { let initialize_uri = &http_request_part.uri; tracing::info!(?initialize_headers, %initialize_uri, "initialize from http server"); } - Ok(self.get_info()) + context.peer.set_peer_info(request.clone()); + self.negotiate_initialize(&request) } } diff --git a/examples/servers/src/common/progress_demo.rs b/examples/servers/src/common/progress_demo.rs index 253dbc28e..91f09473e 100644 --- a/examples/servers/src/common/progress_demo.rs +++ b/examples/servers/src/common/progress_demo.rs @@ -75,13 +75,13 @@ impl ProgressDemo { ctx.meta.get_key_value("progressToken") ); let Some((_, progress_token)) = ctx.meta.get_key_value("progressToken") else { - return Err(McpError::internal_error(format!("No progress token"), None)); + return Err(McpError::internal_error("No progress token", None)); }; let Ok(progress_token) = serde_json::from_value::(progress_token.clone()) else { return Err(McpError::internal_error( - format!("Invalid format of the progress token"), + "Invalid format of the progress token", None, )); }; diff --git a/examples/servers/src/completion_stdio.rs b/examples/servers/src/completion_stdio.rs index 812ed31a8..ffaeee524 100644 --- a/examples/servers/src/completion_stdio.rs +++ b/examples/servers/src/completion_stdio.rs @@ -123,23 +123,23 @@ impl SqlQueryServer { .collect(); // If no uppercase letters found, just use first letter - if first_chars.is_empty() && !candidate.is_empty() { - if let Some(first) = candidate.chars().next() { - first_chars.push(first.to_lowercase().next().unwrap_or('\0')); - } + if first_chars.is_empty() + && let Some(first) = candidate.chars().next() + { + first_chars.push(first.to_lowercase().next().unwrap_or('\0')); } } // Special case: if query is 2 chars and we only got 1 char, try matching first 2 letters - if query_chars.len() == 2 && first_chars.len() == 1 { - if let Some(first) = candidate.chars().nth(0) { - if let Some(second) = candidate.chars().nth(1) { - first_chars = vec![ - first.to_lowercase().next().unwrap_or('\0'), - second.to_lowercase().next().unwrap_or('\0'), - ]; - } - } + if query_chars.len() == 2 + && first_chars.len() == 1 + && let Some(first) = candidate.chars().next() + && let Some(second) = candidate.chars().nth(1) + { + first_chars = vec![ + first.to_lowercase().next().unwrap_or('\0'), + second.to_lowercase().next().unwrap_or('\0'), + ]; } if query_chars.len() != first_chars.len() { @@ -193,7 +193,7 @@ impl SqlQueryServer { } } -#[prompt_router] +#[prompt_router(router = "prompt_router")] impl SqlQueryServer { #[prompt(name = "sql_query", description = "Smart SQL query builder")] async fn sql_query( @@ -308,7 +308,7 @@ impl SqlQueryServer { } } -#[prompt_handler] +#[prompt_handler(router = self.prompt_router)] impl ServerHandler for SqlQueryServer { fn get_info(&self) -> ServerInfo { ServerInfo::new( diff --git a/examples/servers/src/complex_auth_streamhttp.rs b/examples/servers/src/complex_auth_streamhttp.rs index 34c4b1584..bf763d384 100644 --- a/examples/servers/src/complex_auth_streamhttp.rs +++ b/examples/servers/src/complex_auth_streamhttp.rs @@ -69,10 +69,10 @@ impl McpOAuthStore { redirect_uri: &str, ) -> Option { let clients = self.clients.read().await; - if let Some(client) = clients.get(client_id) { - if client.redirect_uri.contains(&redirect_uri.to_string()) { - return Some(client.clone()); - } + if let Some(client) = clients.get(client_id) + && client.redirect_uri == redirect_uri + { + return Some(client.clone()); } None } diff --git a/examples/servers/src/elicitation_enum_inference.rs b/examples/servers/src/elicitation_enum_inference.rs index bed5c1db6..50d2844c6 100644 --- a/examples/servers/src/elicitation_enum_inference.rs +++ b/examples/servers/src/elicitation_enum_inference.rs @@ -97,7 +97,7 @@ struct ElicitationEnumFormServer { tool_router: ToolRouter, } -#[tool_router] +#[tool_router(router = tool_router)] impl ElicitationEnumFormServer { pub fn new() -> Self { Self { @@ -153,7 +153,7 @@ impl ElicitationEnumFormServer { } } -#[tool_handler] +#[tool_handler(router = self.tool_router)] impl ServerHandler for ElicitationEnumFormServer { fn get_info(&self) -> ServerInfo { ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) diff --git a/examples/servers/src/elicitation_stdio.rs b/examples/servers/src/elicitation_stdio.rs index d506a9c7f..16b4773b0 100644 --- a/examples/servers/src/elicitation_stdio.rs +++ b/examples/servers/src/elicitation_stdio.rs @@ -64,7 +64,7 @@ impl Default for ElicitationServer { } } -#[tool_router] +#[tool_router(router = tool_router)] impl ElicitationServer { #[tool(description = "Greet user with name collection")] async fn greet_user( @@ -145,7 +145,7 @@ impl ElicitationServer { } } -#[tool_handler] +#[tool_handler(router = self.tool_router)] impl ServerHandler for ElicitationServer { fn get_info(&self) -> ServerInfo { ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) diff --git a/examples/servers/src/prompt_stdio.rs b/examples/servers/src/prompt_stdio.rs index 7b6e28532..6ef24e937 100644 --- a/examples/servers/src/prompt_stdio.rs +++ b/examples/servers/src/prompt_stdio.rs @@ -112,7 +112,7 @@ impl Default for PromptServer { } } -#[prompt_router] +#[prompt_router(router = "prompt_router")] impl PromptServer { /// Simple greeting prompt without parameters #[prompt( @@ -305,17 +305,17 @@ impl PromptServer { ]; // Add tried solutions if any - if let Some(tried) = args.tried_solutions { - if !tried.is_empty() { - messages.push(PromptMessage::new_text( - Role::User, - format!("I've already tried: {}", tried.join(", ")), - )); - messages.push(PromptMessage::new_text( - Role::Assistant, - "I see you've already attempted some solutions. Let me suggest different approaches.", - )); - } + if let Some(tried) = args.tried_solutions + && !tried.is_empty() + { + messages.push(PromptMessage::new_text( + Role::User, + format!("I've already tried: {}", tried.join(", ")), + )); + messages.push(PromptMessage::new_text( + Role::Assistant, + "I see you've already attempted some solutions. Let me suggest different approaches.", + )); } messages.push(PromptMessage::new_text( @@ -361,7 +361,7 @@ impl PromptServer { } } -#[prompt_handler] +#[prompt_handler(router = self.prompt_router)] impl ServerHandler for PromptServer { fn get_info(&self) -> ServerInfo { ServerInfo::new(ServerCapabilities::builder().enable_prompts().build()).with_instructions( diff --git a/examples/servers/src/sampling_stdio.rs b/examples/servers/src/sampling_stdio.rs index be230add0..2be7f5d46 100644 --- a/examples/servers/src/sampling_stdio.rs +++ b/examples/servers/src/sampling_stdio.rs @@ -1,4 +1,7 @@ -#![allow(deprecated)] +#![expect( + deprecated, + reason = "This example demonstrates the deprecated MCP sampling API" +)] use std::sync::Arc; use anyhow::Result; diff --git a/scripts/keycloak-oauth-fixture.sh b/scripts/keycloak-oauth-fixture.sh new file mode 100755 index 000000000..50c9b85b8 --- /dev/null +++ b/scripts/keycloak-oauth-fixture.sh @@ -0,0 +1,56 @@ +#!/usr/bin/env bash +# ============================================================================= +# keycloak-oauth-fixture.sh — Live OAuth fixture for rmcp refresh tests +# +# Starts a throwaway Keycloak and provisions the realm that +# crates/rmcp/tests/test_live_oauth_refresh.rs expects. Those tests are +# #[ignore]d, so nothing here runs in CI. +# +# Provision: ./scripts/keycloak-oauth-fixture.sh +# Run tests: cargo test -p rmcp --all-features \ +# --test test_live_oauth_refresh -- --ignored +# Tear down: docker rm -f kc-rmcp-test +# +# Requires Docker. Override the port with KC_BASE (default localhost:8081); +# the tests read the same variable. +# ============================================================================= +set -euo pipefail + +KC=${KC_BASE:-http://localhost:8081} +PORT=${KC##*:} +CONTAINER=kc-rmcp-test + +docker rm -f "$CONTAINER" >/dev/null 2>&1 || true +docker run -d --name "$CONTAINER" -p "$PORT:8080" \ + -e KC_BOOTSTRAP_ADMIN_USERNAME=admin -e KC_BOOTSTRAP_ADMIN_PASSWORD=admin \ + quay.io/keycloak/keycloak:26.0 start-dev >/dev/null + +echo "waiting for keycloak at $KC ..." +until curl -sf -o /dev/null "$KC/realms/master/.well-known/openid-configuration"; do + sleep 3 +done + +token=$(curl -s -X POST "$KC/realms/master/protocol/openid-connect/token" \ + -d client_id=admin-cli -d username=admin -d password=admin -d grant_type=password | + python3 -c 'import sys,json;print(json.load(sys.stdin)["access_token"])') + +provision() { + curl -s -X POST "$KC/admin/realms$1" \ + -H "Authorization: Bearer $token" -H "Content-Type: application/json" \ + -d "$2" -o /dev/null -w " ${1:-/} -> %{http_code}\n" +} + +# revokeRefreshToken makes refresh tokens single-use. The concurrency test +# depends on it: without the refresh guard the second caller replays a consumed +# token and Keycloak answers invalid_grant. +provision "" '{"realm":"rmcp","enabled":true,"revokeRefreshToken":true,"refreshTokenMaxReuse":0}' +provision "/rmcp/clients" '{"clientId":"rmcp-client","secret":"rmcp-secret","publicClient":false, + "directAccessGrantsEnabled":true,"standardFlowEnabled":true, + "redirectUris":["http://localhost/callback"]}' +# The profile fields and empty requiredActions keep the direct access grant from +# failing with "Account is not fully set up". +provision "/rmcp/users" '{"username":"alice","enabled":true,"emailVerified":true, + "email":"alice@example.com","firstName":"Alice","lastName":"Example","requiredActions":[], + "credentials":[{"type":"password","value":"alice-pw","temporary":false}]}' + +echo "ready"