diff --git a/.github/ISSUE_TEMPLATE/bug_report.yml b/.github/ISSUE_TEMPLATE/bug_report.yml new file mode 100644 index 0000000..35e2ae6 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/bug_report.yml @@ -0,0 +1,58 @@ +name: Bug report +description: Report incorrect SDK behavior or interoperability +title: "[Bug]: " +body: + - type: markdown + attributes: + value: | + For a security vulnerability, stop here and use the private process in SECURITY.md. + - type: textarea + id: summary + attributes: + label: Summary + description: What happened, and what did you expect instead? + validations: + required: true + - type: input + id: sdk-version + attributes: + label: SDK version or commit + placeholder: v0.2.0 or a full commit SHA + validations: + required: true + - type: input + id: protocol-version + attributes: + label: MCP protocol revision + placeholder: "2025-11-25" + validations: + required: true + - type: textarea + id: environment + attributes: + label: Environment + description: OS, compiler and version, build type, transport, and dependency versions. + validations: + required: true + - type: textarea + id: reproduction + attributes: + label: Minimal reproduction + description: Include the smallest source, build command, and request sequence that reproduces the issue. + validations: + required: true + - type: textarea + id: logs + attributes: + label: Logs or wire messages + description: Remove credentials, tokens, personal data, and other secrets. + render: text + - type: checkboxes + id: checks + attributes: + label: Checklist + options: + - label: I searched existing issues for this problem. + required: true + - label: I removed secrets and sensitive data from the report. + required: true diff --git a/.github/ISSUE_TEMPLATE/config.yml b/.github/ISSUE_TEMPLATE/config.yml new file mode 100644 index 0000000..6a2ce2c --- /dev/null +++ b/.github/ISSUE_TEMPLATE/config.yml @@ -0,0 +1,5 @@ +blank_issues_enabled: false +contact_links: + - name: Security vulnerability + url: https://github.com/yurirocha15/mcp-cpp-sdk/security/policy + about: Report vulnerabilities privately; do not open a public issue. diff --git a/.github/ISSUE_TEMPLATE/feature_request.yml b/.github/ISSUE_TEMPLATE/feature_request.yml new file mode 100644 index 0000000..fb22fb3 --- /dev/null +++ b/.github/ISSUE_TEMPLATE/feature_request.yml @@ -0,0 +1,42 @@ +name: Feature request +description: Propose an SDK capability or usability improvement +title: "[Feature]: " +body: + - type: textarea + id: problem + attributes: + label: Problem + description: What user or interoperability problem should this solve? + validations: + required: true + - type: textarea + id: proposal + attributes: + label: Proposed behavior + description: Describe the public API and observable behavior you would expect. + validations: + required: true + - type: dropdown + id: area + attributes: + label: Area + options: + - Client + - Server + - Protocol types + - Transport + - Authentication or security + - Packaging or build + - Documentation or examples + validations: + required: true + - type: input + id: specification + attributes: + label: Specification or SEP + description: Link the relevant MCP specification or SEP when applicable. + - type: textarea + id: alternatives + attributes: + label: Alternatives + description: What workarounds or alternative designs have you considered? diff --git a/.github/ISSUE_TEMPLATE/question.yml b/.github/ISSUE_TEMPLATE/question.yml new file mode 100644 index 0000000..1d6be7a --- /dev/null +++ b/.github/ISSUE_TEMPLATE/question.yml @@ -0,0 +1,25 @@ +name: Question +description: Ask about SDK usage, behavior, or design +title: "[Question]: " +body: + - type: textarea + id: question + attributes: + label: Question + description: Include the goal you are trying to accomplish. + validations: + required: true + - type: textarea + id: context + attributes: + label: Context + description: Include relevant SDK version, protocol revision, transport, and a concise code sample. + validations: + required: true + - type: checkboxes + id: checks + attributes: + label: Checklist + options: + - label: I checked the README, generated documentation, and existing issues. + required: true diff --git a/.github/labels.yml b/.github/labels.yml new file mode 100644 index 0000000..7978fa2 --- /dev/null +++ b/.github/labels.yml @@ -0,0 +1,36 @@ +- name: bug + color: d73a4a + description: Confirmed or reported defect +- name: enhancement + color: a2eeef + description: New feature or improvement +- name: question + color: d876e3 + description: Usage or design question +- name: needs confirmation + color: fbca04 + description: Needs maintainer confirmation before classification +- name: needs repro + color: f9d0c4 + description: Needs a minimal reproducible example +- name: ready for work + color: 0e8a16 + description: Triaged and ready to implement +- name: good first issue + color: 7057ff + description: Suitable for a first contribution +- name: help wanted + color: 008672 + description: Maintainers welcome community help +- name: P0 + color: b60205 + description: Critical; resolve within 14 calendar days +- name: P1 + color: d93f0b + description: High priority and severe impact +- name: P2 + color: fbca04 + description: Normal project priority +- name: P3 + color: c5def5 + description: Low priority or long-term work diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d9639e3..a691ecb 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -1,6 +1,9 @@ name: CI on: + # Manual runs only. The automatic triggers below stay scoped to `main` on purpose; this + # lets a feature branch be checked on the full matrix without spending CI on every push. + workflow_dispatch: push: branches: [ main ] pull_request: @@ -69,16 +72,42 @@ jobs: restore-keys: ${{ matrix.os }}-cxx${{ matrix.cppstd }}-${{ matrix.linkage }}-ccache- max-size: 500M + # MCP_REQUIRE_TWIN_LOOPBACK turns the OAuthSetterPairAtomicity skips into failures. Those tests + # need a port free on both 127.0.0.1 and 127.0.0.2, and skip when the second address is + # missing -- which would delete the whole setter-pair atomicity regression suite while the run + # still reported green. Linux binds 127.0.0.0/8 as a whole, so there the address is always + # there and its absence is a broken runner, not an unsupported host. macOS gives lo0 only + # 127.0.0.1 unless an alias is added, so the skip stays available there. - name: Build and Test + env: + MCP_REQUIRE_TWIN_LOOPBACK: ${{ runner.os == 'Linux' && '1' || '' }} run: >- python scripts/build.py --examples --test --cppstd ${{ matrix.cppstd }} --linkage ${{ matrix.linkage }} + # The README is the first code a new user copies, so compile it the way they + # would. Skipped on Windows: the check drives the compiler from + # compile_commands.json with GCC/Clang flag spellings. + - name: Compile README snippets + if: runner.os != 'Windows' + run: python scripts/check_readme_snippets.py + - name: Run Examples run: python scripts/run_examples.py + # On the stdio transport stdout is the protocol channel. Running an example + # and checking its exit code does not notice prose printed there, so assert + # the stream separately. run_examples.py builds into build/release. + - name: Check stdio protocol streams + run: python scripts/check_stdio_streams.py build/release + sanitize: + name: sanitize (${{ matrix.compiler }}) runs-on: ubuntu-24.04 + strategy: + fail-fast: false + matrix: + compiler: [default, clang] steps: - uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6 @@ -108,12 +137,108 @@ jobs: - name: Setup ccache uses: hendrikmuhs/ccache-action@e42e6681d2906409c5dde4a315af6214eaa890ee # v1.2 with: - key: ubuntu-24.04-sanitize-ccache-${{ github.ref }} - restore-keys: ubuntu-24.04-sanitize-ccache- + key: ubuntu-24.04-sanitize-${{ matrix.compiler }}-ccache-${{ github.ref }} + restore-keys: ubuntu-24.04-sanitize-${{ matrix.compiler }}-ccache- max-size: 500M - name: Build and Test with Sanitizers - run: python scripts/build.py --sanitize --test + env: + MCP_REQUIRE_TWIN_LOOPBACK: '1' + run: python scripts/build.py --sanitize --compiler ${{ matrix.compiler }} --test + + tsan: + name: tsan (${{ matrix.compiler }}) + runs-on: ubuntu-24.04 + strategy: + fail-fast: false + matrix: + compiler: [default, clang] + + steps: + - uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 + with: + python-version: '3.12' + + - name: Install CMake + uses: jwlawson/actions-setup-cmake@0d6a7d60b009d01c9e7523be22153ff8f19460d3 # v2 + with: + cmake-version: '3.25' + + - name: Cache Conan packages + uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 + with: + path: ~/.conan2/p + key: ubuntu-24.04-conan-tsan-${{ hashFiles('conanfile.py', 'conanfile.txt') }} + restore-keys: ubuntu-24.04-conan- + + - name: Install Dependencies + run: python scripts/init.py + + - name: Setup ccache + uses: hendrikmuhs/ccache-action@e42e6681d2906409c5dde4a315af6214eaa890ee # v1.2 + with: + key: ubuntu-24.04-tsan-${{ matrix.compiler }}-ccache-${{ github.ref }} + restore-keys: ubuntu-24.04-tsan-${{ matrix.compiler }}-ccache- + max-size: 500M + + - name: Build and Test with ThreadSanitizer + env: + MCP_REQUIRE_TWIN_LOOPBACK: '1' + run: python scripts/build.py --tsan --compiler ${{ matrix.compiler }} --test + + msan: + runs-on: ubuntu-24.04 + + steps: + - uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 + with: + python-version: '3.12' + + - name: Install CMake + uses: jwlawson/actions-setup-cmake@0d6a7d60b009d01c9e7523be22153ff8f19460d3 # v2 + with: + cmake-version: '3.25' + + # The dependencies are rebuilt instrumented, so this cache shares nothing with the other jobs. + - name: Cache Conan packages + uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 + with: + path: ~/.conan2/p + key: ubuntu-24.04-conan-msan-${{ hashFiles('conanfile.py', 'conanfile.txt') }} + restore-keys: ubuntu-24.04-conan-msan- + + # Keyed on MSAN_LLVM_COMMIT in scripts/build.py. A cache from another commit is rebuilt by + # build.py rather than used, so a key left behind costs time and nothing else. + - name: Cache instrumented libc++ + uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 + with: + path: build/msan-libcxx + key: ubuntu-24.04-msan-libcxx-3b5b5c1ec4a3 + + - name: Install Dependencies + run: python scripts/init.py + + - name: Setup ccache + uses: hendrikmuhs/ccache-action@e42e6681d2906409c5dde4a315af6214eaa890ee # v1.2 + with: + key: ubuntu-24.04-msan-ccache-${{ github.ref }} + restore-keys: ubuntu-24.04-msan-ccache- + max-size: 500M + + - name: Build and Test with MemorySanitizer + env: + MCP_REQUIRE_TWIN_LOOPBACK: '1' + run: python scripts/build.py --msan --test release-safety: name: Release safety contracts @@ -134,7 +259,8 @@ jobs: python3 scripts/check_release_workflow.py PYTHONPATH=scripts python3 -m unittest -v \ scripts.test_release_dispatch_contract \ - scripts.test_release_workflow_policy + scripts.test_release_workflow_policy \ + scripts.test_check_json_matrix - name: Validate release artifacts and package templates run: | diff --git a/.github/workflows/conformance-guards.yml b/.github/workflows/conformance-guards.yml new file mode 100644 index 0000000..1df63b2 --- /dev/null +++ b/.github/workflows/conformance-guards.yml @@ -0,0 +1,32 @@ +name: Conformance guardrails + +on: + # Manual runs only. The automatic triggers below stay scoped to `main` on purpose; this + # lets a feature branch be checked on the full matrix without spending CI on every push. + workflow_dispatch: + push: + branches: [main] + pull_request: + branches: [main] + +permissions: + contents: read + +jobs: + alpha-runner-non-gating: + name: Alpha runner stays non-gating + runs-on: ubuntu-24.04 + timeout-minutes: 5 + + steps: + - uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 + with: + python-version: '3.12' + + - name: Assert the alpha conformance runner cannot gate CI + run: python3 scripts/check_alpha_runner_guard.py diff --git a/.github/workflows/conformance.yml b/.github/workflows/conformance.yml new file mode 100644 index 0000000..f58538b --- /dev/null +++ b/.github/workflows/conformance.yml @@ -0,0 +1,74 @@ +name: MCP conformance regression baseline + +on: + # Manual runs only. The automatic triggers below stay scoped to `main` on purpose; this + # lets a feature branch be checked on the full matrix without spending CI on every push. + workflow_dispatch: + push: + branches: [main] + pull_request: + branches: [main] + +permissions: + contents: read + +jobs: + conformance-2025-11-25: + name: Pinned client/server baseline (2025-11-25) + runs-on: ubuntu-24.04 + timeout-minutes: 45 + + steps: + - uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 + with: + python-version: '3.12' + + - name: Set up Node.js + uses: actions/setup-node@6044e13b5dc448c55e2357c09f80417699197238 # v6.2.0 + with: + node-version: '22' + cache: npm + cache-dependency-path: conformance/runner/package-lock.json + + - name: Set up ccache + uses: hendrikmuhs/ccache-action@e42e6681d2906409c5dde4a315af6214eaa890ee # v1.2 + with: + key: ubuntu-24.04-conformance + + - name: Install CMake + uses: jwlawson/actions-setup-cmake@0d6a7d60b009d01c9e7523be22153ff8f19460d3 # v2 + with: + cmake-version: '3.25' + + - name: Cache Conan packages + uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5 + with: + path: ~/.conan2/p + key: ubuntu-24.04-conan-conformance-${{ hashFiles('conanfile.py', 'conanfile.txt') }} + restore-keys: ubuntu-24.04-conan- + + - name: Install C++ dependencies + run: python3 scripts/init.py + + - name: Build conformance fixtures + run: python3 scripts/build.py --conformance --jobs "$(nproc)" + + - name: Install locked official runner + run: npm ci --ignore-scripts --prefix conformance/runner + + - name: Run pinned regression suites + run: bash conformance/run.sh build/release build/conformance-results + + - name: Upload conformance evidence + if: always() + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 + with: + name: mcp-conformance-2025-11-25 + path: build/conformance-results + if-no-files-found: error + retention-days: 30 diff --git a/.github/workflows/docs.yml b/.github/workflows/docs.yml index f66bb03..4b898b9 100644 --- a/.github/workflows/docs.yml +++ b/.github/workflows/docs.yml @@ -1,6 +1,9 @@ name: Docs on: + # Manual runs only. The automatic triggers below stay scoped to `main` on purpose; this + # lets a feature branch be checked on the full matrix without spending CI on every push. + workflow_dispatch: push: branches: [main] diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 66924ff..cf6dedd 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -1497,16 +1497,28 @@ jobs: id: publish env: AUR_SSH_PRIVATE_KEY_B64: ${{ secrets.AUR_SSH_PRIVATE_KEY_B64 }} + AUR_SSH_KEY_PASSPHRASE: ${{ secrets.AUR_SSH_KEY_PASSPHRASE }} AUR_KNOWN_HOSTS: ${{ vars.AUR_KNOWN_HOSTS }} VERSION: ${{ needs.contract.outputs.version }} TAG: ${{ inputs.tag }} run: | install -d -m 700 "${HOME}/.ssh" - trap 'rm -f "${HOME}/.ssh/id_ed25519"' EXIT + umask 077 + askpass="${RUNNER_TEMP}/aur-ssh-askpass" + trap 'ssh-agent -k >/dev/null 2>&1 || true; rm -f "${HOME}/.ssh/id_ed25519" "${HOME}/.ssh/known_hosts" "${askpass}"' EXIT printf '%s\n' "${AUR_KNOWN_HOSTS}" > "${HOME}/.ssh/known_hosts" grep -E '^aur\.archlinux\.org[ ,]' "${HOME}/.ssh/known_hosts" >/dev/null printf '%s' "${AUR_SSH_PRIVATE_KEY_B64}" | base64 --decode > "${HOME}/.ssh/id_ed25519" chmod 600 "${HOME}/.ssh/id_ed25519" + test -n "${AUR_SSH_KEY_PASSPHRASE}" + printf '%s\n' '#!/bin/sh' 'printf "%s\\n" "${AUR_SSH_KEY_PASSPHRASE}"' > "${askpass}" + chmod 700 "${askpass}" + export SSH_ASKPASS="${askpass}" SSH_ASKPASS_REQUIRE=force DISPLAY=release + eval "$(ssh-agent -s)" + ssh-add "${HOME}/.ssh/id_ed25519" ` gains `BearerChallengeConfig`, + `ProtectedResourceMetadataConfig`, `format_www_authenticate()`, + `protected_resource_metadata_path()`, `protected_resource_metadata_url()`, + `format_protected_resource_metadata()` and `http_request_path()`. +- `set_async_bearer_token_validator()` on `HttpServerTransport` and + `StreamableHttpSessionManager`, taking the new `AsyncBearerTokenValidator` + (`std::function(std::string)>`), so a token decision that needs + I/O — introspection, a JWKS fetch — suspends instead of blocking the + executor serving MCP traffic. Installing it alongside the synchronous + validator throws `std::logic_error`. +- `set_max_request_body_bytes()` on `HttpServerTransport` and + `StreamableHttpSessionManager`, together with + `mcp::constants::g_default_max_request_body_bytes` (8 MiB), the new default. + A body over the cap is answered `413 Payload Too Large` and its connection + closed before MCP dispatch; zero is rejected with `std::invalid_argument`. + The default replaces Beast's own 1 MB limit, which admitted only about + 750 KB of raw content once base64 inflation inside the JSON body is + accounted for. +- `ClientOptions::on_protocol_error`, invoked when the client discards an + incoming message instead of dispatching it: a message the peer sent that + could not be decoded, reported as `g_PARSE_ERROR`, or an exception thrown by + an application notification callback, reported as `g_INTERNAL_ERROR` with + the notification method in the message. It runs on the read loop and must + not block. +- Documentation for the server-side OAuth challenge work: the OAuth guide now + names the protected-resource metadata helpers and the exported bearer and + request-path helpers, and renders its example challenge through the + server-side API rather than hand-writing the header. The client's default + 30-second request timeout is documented for the first time, including that + the typed helpers such as `call_tool()` and `read_resource()` accept no + per-request override, that a request outliving it fails with + `g_REQUEST_TIMEOUT` while the peer may still be running it, and that no + value disables the deadline. +- `python scripts/build.py --tsan --test` builds and runs the tests under + ThreadSanitizer (CMake option `ENABLE_TSAN`), and `--msan --test` under + MemorySanitizer, against a libc++ and dependencies it builds instrumented. + `--compiler clang` compiles the AddressSanitizer and ThreadSanitizer builds + with Clang. CI runs all of them on every push, and the AddressSanitizer run + now also checks for stack use after return. +- `WebSocketClientTransport` takes a `connect_timeout` (default 30 seconds, + zero disables it) that bounds the TCP connect and the WebSocket handshake + together. A peer that accepts the connection and then stalls now fails the + pending call with a timeout instead of holding it forever. + +### Changed + +- Implementation moved out of oversized headers into compiled translation + units (OAuth, client runtime, protocol tools, memory transport, HTTP + types); protocol models and typed handler templates remain header-based. +- Documentation guides and feature examples refreshed to match the compiled + runtime split. +- OAuth protected-resource metadata that omits the RFC 9728 `resource` + member is now rejected instead of falling back to the configured server + URL; a `resource` value must identify the configured server (exact match + or an origin/segment-boundary prefix). +- `OAuthHttpClient` constructed without a `MetadataFetchPolicy` now refuses + every request (deny-all default) instead of allowing any target; callers + must supply an explicit policy. +- A `MetadataFetchPolicy::denied_origins` entry that is not a bare origin — it + carries a path (a lone trailing `/` included), a query or a fragment, or its + port is not a plain in-range decimal number — is now rejected instead of + silently ignored. `validate_metadata_url` throws `MetadataPolicyError` with + the new `MetadataUrlDecision::denied_origin_entry_malformed`, naming the + offending entry, and refuses every target until the policy is corrected. + Previously such an entry denied nothing, so `denied_origins` of + `{"https://evil.example/"}` admitted `https://evil.example`. A configuration + that relied on that silence is now a hard failure. `allowed_origins` is + unchanged: an entry that is not a bare origin still matches nothing, because + dropping an allow entry grants nothing and so fails closed. +- Injected OAuth client credentials that carry a `client_secret` must now name + the authorization server they are bound to, via the new + `OAuthAuthorizationConfig::client_issuer` or + `ClientIdentityConfig::pre_registered`'s `issuer`. A secret that names no + issuer is refused at the point of use: `select_client_identity` returns + `ClientIdentityDecision::unavailable` and the authorization attempt fails + with a message naming the missing binding, rather than presenting the secret. + Previously an empty `issuer` fell through to `use_pre_registered` for every + authorization server, so the existing misbinding guard was inert on both + paths the SDK itself constructs, and a hostile-but-policy-allowed + authorization server named in a protected-resource document received the + application's secret at its token endpoint. "Bound to no issuer" is not + "bound to every issuer". Construction is unchanged, so a caller breaks only + when it actually attempts the affected flow; a public client (a `client_id` + with no secret) is unaffected and still authorizes against any issuer. +- An optional member serialized as an explicit `null` is now read the same way + as an absent one. The protocol `from_json` overloads test presence through + the new `detail::has_json_value` helper instead of `contains()`, which + previously accepted the null and then threw when the value was extracted; + this covers `Error::data`, `RelatedTaskMetadata::title`, `TaskMetadata`'s + `ttl` and `relatedTasks`, request and notification `params`, response `id` + and `error`, and the notification `_meta`, `reason`, `total`, `message`, + `logger` and `metadata` members. An `Error` whose `message` is missing or + null now decodes to an empty message rather than throwing, so the `code` a + caller acts on survives. The server reads a null `error` member as absent + when validating request, notification and response envelopes and when + dispatching a response, and a null notification `params` as no params. The + client is asymmetric on purpose: a null `error` is the absence of an error, + while a null `result` is a legitimate empty result and still counts as + present. The same treatment covers the optional members of the protocol + types a peer's payload reaches: `description`, `mimeType`, `size`, `title` + and `icons` on `Resource` and `ResourceTemplate`; `description`, `title`, + `icons` and `execution` on `Tool`; and `isError` on both `CallToolResult` + and `ToolResultContent`. +- A peer whose serializer writes absent optionals as explicit nulls could not + complete initialization at all. `clientInfo` is an `Implementation`, whose + `from_json` guarded `title`, `description`, `websiteUrl` and `icons` with a + bare presence test, so `"clientInfo": {"name": "x", "version": "1", + "title": null}` failed to decode and the `initialize` request was answered + `-32602`. Because `initialize` is the peer's first message, the failure was + unconditional: no such client could connect at all. Those members are now + read as absent when they arrive as null. +- Explicit nulls in request parameters that previously drew `-32602` are now + read as absent: `arguments` on `prompts/get`, `context.arguments` on + `completion/complete`, and the pagination `cursor`, where `"cursor": null` + on a list request was rejected as an invalid cursor instead of returning the + first page. `CallToolParams::arguments` is deliberately unchanged: it is a + required member carrying a shape rule rather than an optional one, so + rejecting a null there remains correct. +- An explicit null `_meta` or `annotations` no longer decodes to a value the + peer never sent. Unlike the members above these never threw: the guard + accepted the null and produced an engaged optional, so the SDK re-emitted + `"_meta": null` and turned `"annotations": null` into `"annotations": {}` — + a structurally present `Annotations` that a consumer reads as "annotations + supplied, carrying no constraints" rather than as none at all. Neither side + saw anything wrong, which is what made this worse than a rejection. The + affected members are `_meta` and `annotations` on `Resource`, + `ResourceTemplate`, `Tool`, `CallToolResult` and `ToolResultContent`, + `Tool::outputSchema`, `structuredContent` on `CallToolResult` and + `ToolResultContent`, and `CallToolParams::_meta` on the `tools/call` path. +- A `resources/read` URI longer than 512 characters is now answered + `-32602` naming the limit, instead of being handed to the regular-expression + template matcher, whose stack use grows with the subject and can overrun the + smallest thread stack the SDK runs on. Exact resource lookups happen first + and are unaffected by the limit. +- A resource URI template with two expressions and no literal between them + (`{a}{b}`) is now rejected at registration with `std::invalid_argument`. + Such a template compiled to adjacent unbounded runs that the matcher could + only resolve by backtracking over every split of the input, and the boundary + between the two variables is undecidable in any case. +- An exception from a tool handler is reported as a `CallToolResult` carrying + `isError`, with the message sanitized before it reaches the peer, and the + guard now covers a throw that does not derive from `std::exception`. Such a + throw previously escaped every guard on the request path — the tool + invocation caught only `std::exception`, and `dispatch_request_wire` has no + catch-all — so no response was written at all: a session client waited out + its own request timeout, and a stateless connection was closed with nothing + sent. It is now reported as a tool error with the fixed text + `Tool handler failed`, because such a throw carries no message the SDK can + quote. This is the behavior the error-handling guide already documented for + handler exceptions, "typed, asynchronous or raw". The change is specific to + tool handlers: a throw that does not derive from `std::exception` escaping a + resource or prompt handler, or middleware, still drops the response. +- Middleware now runs outside the tool handler's exception guard, so a + middleware that throws surfaces as a JSON-RPC error (`-32603`) rather than + as a tool error result; middleware decides whether a call may proceed at + all, which is a protocol-level answer rather than a tool outcome. Previously + any exception from the middleware chain became a tool error result. The + bundled auth middleware no longer throws on a missing or invalid bearer + token — it returns a tool error result — so that rejection still reaches + the caller as a tool result rather than becoming `-32603`. +- A handler result that is already a serialized `CallToolResult` is now passed + through unwrapped instead of being nested inside a text block. Domain data + that merely carries a `content` key fails the full `CallToolResult` check + and is still wrapped. +- An incoming message the client cannot decode no longer ends the session. + Only a failure of `read_message()` — the peer hung up, the socket died, the + client was closed — fails pending requests and closes the transport; + anything that goes wrong after the bytes are off the wire is reported + through `ClientOptions::on_protocol_error` and the message is dropped. +- An explicit `ProtectedResourceMetadataConfig::path` is now validated rather + than concatenated onto the origin verbatim. A path that does not begin with + `/` is rejected with `std::invalid_argument`, as is one carrying a `.` or + `..` whole segment. Previously `path = "evil"` produced `https://hevil` — a + corrupted authority rather than merely a bad path — and `path = "/../../x"` + was published unresolved. Only whole segments count, so a path such as + `/.well-known/a..b/c.d` is still accepted; dots inside a segment are + ordinary characters, and refusing them would refuse the well-known prefix + the derivation itself produces. The function does not normalize, and its + result is published to clients as the authoritative location of the + document, so an unusable path is refused rather than advertised. The + transport and session-manager setters reject it at configuration time and + leave the server unchanged. +- A `Server` may now be destroyed while handlers are still in flight. The + implementation outlives the `Server` until that work finishes, so a handler + no longer runs against freed state, and a reverse request issued from a + handler whose `Server` is gone throws `std::runtime_error` naming that cause + instead of faulting — distinct from the failure reported for a session that + is merely closing. Destroying a `Server` from a thread other than the one + running its session no longer races on that session's request maps either: + those entries are abandoned on the session strand, which is the only + executor permitted to touch them. +- A session is now unregistered even when teardown throws — a transport that + throws from `close()`, or a drain wait that fails for any reason other than + cancellation. Previously such a failure left the session registered and + every later `run()` was refused for the lifetime of the `Server`. +- Stateless and discover dispatch in the HTTP session manager is spelled as a + named coroutine rather than a lambda handed to `co_spawn`, so the request + moves into that coroutine's own frame. This avoids the GCC 11 defect that + corrupts an object living in a coroutine frame across a suspension. +- The peer-input matrix build check now fails when + `test/core/json_peer_input_matrix_test.cpp` differs from what + `scripts/gen_json_matrix.py` generates, and exits 2 for a row whose decoder + now throws on an explicit null or an absent key. The generator refuses to + write such a row unless it is named with `--accept-regression`. The check is + governed by the new `MCP_CPP_SDK_CHECK_JSON_MATRIX` option: it defaults to + `ON` in a checkout, where a `BUILD_TESTING=ON` build now needs Python 3.9+, + and to `OFF` in source archives, which do not ship the scripts. Previously a + missing Python skipped the check with a warning, and a source-archive build + with tests failed because the scripts were absent. +- Closing an OAuth client transport from a thread that does not run the + `io_context` could be lost while a metadata lookup was resolving: the close + found no socket to shut, and the flow went on to connect and blocked until + the HTTP timeout. The abort latch is now checked again just before the + connection is opened. `OAuthHttpClient` also keeps every request on the + client's one strand, so this holds when several threads run the `io_context`. + A request used to leave that strand at its first suspension, so a close could + run on another thread at the same time: a data race on the socket that could + crash, or a close lost in the instant before the socket opened. Callers are + still resumed on their own executor and never on the client's strand, + including a caller that has been cancelled. +- Closing `HttpClientTransport` or `WebSocketClientTransport` from a thread + that does not run the `io_context` could be lost while the transport was + resolving its server's address: the close found no socket to act on, and the + transport went on to connect. The HTTP write blocked until the HTTP timeout, + and the WebSocket connect waited on its handshake with no limit. Both + transports now check for the close again once the address is known and fail + the pending call instead of connecting. For `WebSocketClientTransport` this + holds however many threads run the `io_context`. +- `HttpClientTransport::close()` now closes the socket of a write that is in + flight instead of cancelling its pending operation once. A close that arrived + between two socket operations, or just after one had completed, cancelled + nothing and the write carried on until the HTTP timeout; with several threads + running the `io_context` that included a close racing the connect. The write + now fails promptly in each of these cases, however many threads run the + `io_context`, and reports `operation_aborted` on every platform, whatever the + closed socket itself reported. The session `DELETE` that `close()` sends + afterwards uses a new connection. It already did after a cancelled write; + when the response had already arrived as `close()` ran, it used to reuse the + write's connection and now opens a new one as well. + ## [0.2.0] - TBD ### Added diff --git a/CMakeLists.txt b/CMakeLists.txt index 268af8e..f9fb3d2 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -85,11 +85,21 @@ configure_file( set(MCP_CPP_SDK_SOURCES src/core/version.cpp src/core/runtime.cpp + src/core/secure_random.cpp + src/core/serialized_transport_writer.cpp + src/protocol/tools.cpp + src/auth/challenge.cpp + src/auth/client_identity.cpp + src/auth/metadata_policy.cpp + src/auth/oauth.cpp + src/client/client.cpp src/server/server_stdio.cpp src/server/server_http.cpp src/server/server_tool.cpp src/server/server.cpp src/transport/stdio.cpp + src/transport/memory.cpp + src/transport/http_types.cpp src/transport/http_server.cpp src/transport/http_client.cpp src/transport/http_session_manager.cpp @@ -155,6 +165,14 @@ target_link_libraries(mcp-cpp-sdk INTERFACE "mcp-cpp-sdk-${MCP_CPP_SDK_DEFAULT_LINKAGE}") add_library(mcp::sdk ALIAS mcp-cpp-sdk) +option( + MCP_CPP_SDK_BUILD_CONFORMANCE + "Build client and server fixtures for the official MCP conformance runner" + OFF) +if(MCP_CPP_SDK_BUILD_CONFORMANCE) + add_subdirectory(conformance) +endif() + # Documentation target option(BUILD_DOCS "Add the Doxygen documentation target" ON) if(BUILD_DOCS) @@ -174,6 +192,25 @@ endif() # Build Tests option(BUILD_TESTING "Build tests" ON) + +# ENABLE_SANITIZERS is declared and applied inside the block below, so with +# testing off the flag is accepted and does nothing: the build lands in a +# directory named for sanitizers without being sanitized. Refuse rather than +# hand back a clean run that proved nothing. +if(ENABLE_SANITIZERS AND NOT BUILD_TESTING) + message( + FATAL_ERROR + "ENABLE_SANITIZERS requires BUILD_TESTING=ON; with tests off the " + "sanitizer flags are never applied and the build is not sanitized") +endif() + +if(ENABLE_TSAN AND NOT BUILD_TESTING) + message( + FATAL_ERROR + "ENABLE_TSAN requires BUILD_TESTING=ON; with tests off the " + "sanitizer flags are never applied and the build is not sanitized") +endif() + if(BUILD_TESTING) find_package(GTest REQUIRED) enable_testing() @@ -227,14 +264,20 @@ if(BUILD_TESTING) -P "${MCP_CPP_SDK_PC_TEST_SCRIPT}") endif() set(TEST_SRCS + test/auth/auth_authorization_test.cpp + test/auth/auth_challenge_test.cpp + test/auth/auth_client_identity_test.cpp test/auth/auth_integration_test.cpp test/auth/auth_oauth_test.cpp + test/auth/auth_token_endpoint_auth_test.cpp test/client/client_core_test.cpp test/client/client_features_test.cpp test/client/client_notifications_test.cpp + test/core/capabilities_extensions_test.cpp test/core/concepts_test.cpp test/core/core_test.cpp test/core/protocol_test.cpp + test/core/serialized_transport_writer_test.cpp test/server/roots_test.cpp test/server/sampling_test.cpp test/server/server_core_test.cpp @@ -259,12 +302,81 @@ if(BUILD_TESTING) test/server/server_http_test.cpp test/server/server_tool_test.cpp test/transport/transport_factory_test.cpp - test/core/gcc11_sso_crash_test.cpp) + test/core/gcc11_sso_crash_test.cpp + test/core/json_peer_input_matrix_test.cpp) add_executable(mcp-sdk-tests ${TEST_SRCS}) target_link_libraries(mcp-sdk-tests PRIVATE mcp-cpp-sdk GTest::gtest_main) + if(CMAKE_SYSTEM_NAME STREQUAL "Linux") + # The resolve and socket gates interpose getaddrinfo() and socket() and + # forward with dlsym(), which lives in libdl before glibc 2.34. + target_sources(mcp-sdk-tests PRIVATE test/support/resolve_gate.cpp + test/support/socket_gate.cpp) + target_link_libraries(mcp-sdk-tests PRIVATE ${CMAKE_DL_LIBS}) + endif() + if(MSVC) + # A gtest translation unit emits a COMDAT section per template + # instantiation, and the suites built on nlohmann::json and Boost.Asio + # instantiate enough of them to pass the 65,279 sections a COFF object can + # hold (error C1128). The flag lifts that limit to 2^32; the only cost is a + # larger object file. + target_compile_options(mcp-sdk-tests PRIVATE /bigobj) + endif() + + # The peer-input matrix is generated from scripts/json_census.py, and its + # worth depends on covering every protocol type. Checking that here means a + # new type with a from_json breaks the build until it is either in the matrix + # or excluded with a stated reason, rather than quietly arriving untested. + # Source archives ship test/ but not the matrix scripts, so a packager's + # BUILD_TESTING=ON build defaults the check off instead of failing on it. + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/scripts/check_json_matrix.py") + set(_mcp_json_matrix_default ON) + else() + set(_mcp_json_matrix_default OFF) + endif() + option( + MCP_CPP_SDK_CHECK_JSON_MATRIX + "Fail the build when the peer-input matrix is stale (needs Python 3.9+)" + ${_mcp_json_matrix_default}) + if(MCP_CPP_SDK_CHECK_JSON_MATRIX) + if(NOT EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/scripts/check_json_matrix.py") + message( + FATAL_ERROR + "MCP_CPP_SDK_CHECK_JSON_MATRIX=ON but scripts/check_json_matrix.py " + "is not in this source tree") + endif() + find_package(Python3 3.9 REQUIRED COMPONENTS Interpreter) + add_custom_target( + mcp-json-matrix-complete ALL + COMMAND "${Python3_EXECUTABLE}" + "${CMAKE_CURRENT_SOURCE_DIR}/scripts/check_json_matrix.py" + WORKING_DIRECTORY "${CMAKE_CURRENT_SOURCE_DIR}" + COMMENT "Checking the peer-input matrix covers every protocol type") + add_dependencies(mcp-sdk-tests mcp-json-matrix-complete) + # Inside the gate for the same reason as the check: source archives do not + # ship these scripts. + add_test( + NAME json-matrix-check-unittest + COMMAND "${Python3_EXECUTABLE}" scripts/test_check_json_matrix.py -v + WORKING_DIRECTORY "${CMAKE_CURRENT_SOURCE_DIR}") + set_tests_properties(json-matrix-check-unittest + PROPERTIES ENVIRONMENT PYTHONDONTWRITEBYTECODE=1) + elseif(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/scripts/check_json_matrix.py") + message(STATUS "Peer-input matrix check skipped " + "(disabled by MCP_CPP_SDK_CHECK_JSON_MATRIX=OFF)") + else() + message( + STATUS + "Peer-input matrix check skipped (MCP_CPP_SDK_CHECK_JSON_MATRIX=OFF; " + "scripts/ is not shipped in source archives)") + endif() include(GoogleTest) - gtest_discover_tests(mcp-sdk-tests) + # Default per-test timeout so one deadlocked test cannot stall a ctest run for + # the 25-minute ctest default (or forever without --timeout). + gtest_discover_tests( + mcp-sdk-tests + PROPERTIES + TIMEOUT 120) option(ENABLE_COVERAGE "Enable coverage reporting" OFF) if(ENABLE_COVERAGE) @@ -276,6 +388,31 @@ if(BUILD_TESTING) option(ENABLE_SANITIZERS "Enable AddressSanitizer + UndefinedBehaviorSanitizer" OFF) + option(ENABLE_TSAN "Enable ThreadSanitizer" OFF) + if(ENABLE_TSAN AND ENABLE_SANITIZERS) + message( + FATAL_ERROR + "ENABLE_TSAN and ENABLE_SANITIZERS cannot be combined: ThreadSanitizer " + "does not run with AddressSanitizer") + endif() + if(ENABLE_TSAN) + if(CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang") + # -g1 keeps the line tables a report needs and drops the rest. + # ThreadSanitizer stops every thread while it symbolizes a report, + # including one a suppression then discards, and that pause grows with the + # debug information it has to read: with full -g it outlasted the tests' + # own time bounds on CI. + set(TSAN_FLAGS -fsanitize=thread -fno-omit-frame-pointer -g1 -O1) + foreach(sdk_target IN LISTS MCP_CPP_SDK_INSTALL_TARGETS) + target_compile_options(${sdk_target} PUBLIC ${TSAN_FLAGS}) + target_link_options(${sdk_target} PUBLIC -fsanitize=thread) + endforeach() + target_compile_options(mcp-sdk-tests PRIVATE ${TSAN_FLAGS}) + target_link_options(mcp-sdk-tests PRIVATE -fsanitize=thread) + else() + message(WARNING "ENABLE_TSAN is only supported with GCC or Clang") + endif() + endif() if(ENABLE_SANITIZERS) if(CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang") set(SANITIZER_FLAGS -fsanitize=address,undefined -fno-omit-frame-pointer diff --git a/DEPENDENCY_POLICY.md b/DEPENDENCY_POLICY.md new file mode 100644 index 0000000..3844d19 --- /dev/null +++ b/DEPENDENCY_POLICY.md @@ -0,0 +1,63 @@ +# Dependency Policy + +This policy covers dependencies required to build, link, test, package, and +release `mcp-cpp-sdk`. + +## Supported runtime dependencies + +The installed SDK has three direct dependencies: + +| Dependency | Supported floor | Purpose | +| --- | --- | --- | +| Boost | 1.74 | Asio executors, coroutines, and networking | +| nlohmann/json | 3.10.5 | JSON and JSON-RPC serialization | +| OpenSSL | 3.0 | Cryptographic randomness and OAuth helper primitives | + +The minimum versions in `CMakeLists.txt` are the source of truth. CI must test +those floors on at least one supported platform before a stable release. The +latest stable dependency versions are tested periodically to detect upcoming +compatibility problems. + +The SDK's built-in HTTP transports currently use plaintext HTTP. OpenSSL is +not used to provide TLS transport; deployments that leave a loopback or trusted +network boundary must use a TLS-terminating proxy or a custom TLS transport. + +Build, test, documentation, and release-only dependencies do not become part +of the public link interface. Their exact versions should be locked where the +tool supports a lock file or pinned by immutable revision in CI. + +## Updates and support windows + +- Patch and minor dependency updates may be adopted in any SDK release when + they preserve the documented compiler, platform, source, and ABI contracts. +- Raising a direct dependency floor is announced in the changelog. Before + `1.0.0` it requires at least a minor SDK release; after `1.0.0` it requires a + major release unless the old dependency is unsupported or has an unresolved + vulnerability. +- A dependency release that is end-of-life or prevents protocol conformance may + be removed from support after notice in the changelog and roadmap. +- Unsupported or unmaintained transitive dependencies are replaced when a + maintained alternative exists and the migration risk is reasonable. + +## Adding dependencies + +A new direct dependency must have a compatible license, active maintenance, +documented security reporting, supported CMake consumption, and a demonstrated +benefit that is not reasonably achievable with the C++ standard library or an +existing dependency. Optional features should keep their dependencies private +and optional whenever possible. + +## Vulnerabilities + +Report vulnerabilities through the private process in `SECURITY.md`. A known +dependency vulnerability is classified using the maintenance priorities in +`MAINTENANCE.md`; CVSS 7.0 or higher is P0. Remediation may include upgrading, +backporting, disabling the affected feature, or documenting that the SDK is not +reachable. Security updates can override the normal compatibility window, and +the release notes must explain any resulting consumer action. + +## Review cadence + +Maintainers review direct dependencies at least monthly, before each release, +and whenever a relevant security advisory is published. Dependency changes are +covered by the normal build, sanitizer, package-consumer, and conformance gates. diff --git a/MAINTENANCE.md b/MAINTENANCE.md new file mode 100644 index 0000000..cbf6a59 --- /dev/null +++ b/MAINTENANCE.md @@ -0,0 +1,47 @@ +# Maintenance Policy + +This project uses public issue labels and milestones to make maintenance status +observable. Security reports follow the private process in `SECURITY.md`. + +This policy is effective prospectively for issues opened on or after +2026-07-19. It does not establish historical response-time evidence. + +## Service levels + +- New public issues are triaged within 30 calendar days. +- P0 issues are resolved, mitigated, or have a safe release available within + 14 calendar days of the initial report. +- The Tier 1 target is triage within two business days and P0 resolution within + seven calendar days; those shorter windows become release policy only after + the project has demonstrated that capacity. + +Triage means reproducing or requesting the information needed to reproduce, +classifying the issue, and identifying the next state. Actionable issues also +receive one priority label. An acknowledgement without classification is not +complete triage. + +## Priority + +- **P0:** CVSS 7.0 or higher, failure of core MCP operations for supported + users, data loss, authentication bypass, or a release-blocking regression. +- **P1:** severe degradation without a reasonable workaround or a major + conformance regression. +- **P2:** normal correctness, interoperability, performance, or usability work. +- **P3:** low-impact cleanup, polish, or long-term improvement. + +## Workflow labels + +Every triaged issue uses one type label: `bug`, `enhancement`, or `question`. +It also receives one of `needs confirmation`, `needs repro`, or `ready for +work`. Actionable issues receive exactly one of `P0` through `P3`. Issues +suitable for outside contributors may additionally use `good first issue` or +`help wanted`. The canonical label definitions are stored in +`.github/labels.yml`; maintainers must apply that manifest to the live +repository before using it as tier evidence. + +## Supported releases + +Before `1.0.0`, only the latest minor release receives fixes. Starting with +`1.0.0`, the latest major release is supported; older major lines receive only +explicitly announced security backports. Supported compilers, platforms, and +dependency floors are those exercised by CI and documented for the release. diff --git a/README.md b/README.md index a5b637f..306176f 100644 --- a/README.md +++ b/README.md @@ -5,14 +5,24 @@ A modern C++20 implementation of the Model Context Protocol (MCP), enabling seam [![C++20](https://img.shields.io/badge/C%2B%2B-20-blue.svg)](https://isocpp.org/) [![License: Apache 2.0](https://img.shields.io/badge/License-Apache%202.0-blue.svg)](LICENSE) [![Build Status](https://github.com/yurirocha15/mcp-cpp-sdk/actions/workflows/ci.yml/badge.svg)](https://github.com/yurirocha15/mcp-cpp-sdk/actions/workflows/ci.yml) +[![Conformance baseline](https://github.com/yurirocha15/mcp-cpp-sdk/actions/workflows/conformance.yml/badge.svg)](https://github.com/yurirocha15/mcp-cpp-sdk/actions/workflows/conformance.yml) ## Why mcp-cpp-sdk? -- **Modern C++20**: Asynchronous first, leveraging coroutines (via Boost.Asio) for high-performance I/O. +- **Modern C++20**: Coroutine-based asynchronous I/O with Boost.Asio. - **Shared or Static**: The same public API is available through explicit CMake targets for either linkage model. -- **Type-Safe Protocol**: Strong typing for all MCP messages using `nlohmann/json`. +- **Typed Protocol Models**: Strongly typed models for the supported MCP surface, with `nlohmann/json` interoperability. - **Flexible Transports**: Native support for Stdio, WebSocket, and Streamable HTTP. -- **Full Specification**: Complete implementation of the latest MCP protocol (2025-11-25). +- **Measured Interoperability**: Official client and server conformance suites run in CI against a pinned `2025-11-25` regression baseline, with unsupported scenarios kept visible. A green baseline means no drift, not that an SDK tier has been achieved. + +### How this implementation differs + +The SDK combines a compiled client/server/transport runtime with header-based +protocol models and typed handler templates. Ordinary tool return values are +normalized into protocol-valid structured results, while complete raw +`CallToolResult` values use an explicit API. Interoperability and performance +claims are kept reproducible through the pinned conformance baseline and the +shared-workload benchmark adapters in [`benchmark/`](benchmark/). ## Quick Start @@ -41,15 +51,18 @@ target_link_libraries(your_target PRIVATE mcp::sdk) #include int main() { - mcp::Server server({"hello-server", "1.0.0"}, {}); + mcp::ServerCapabilities capabilities; + capabilities.tools = mcp::ServerCapabilities::ToolsCapability{}; + mcp::Server server({"hello-server", "1.0.0"}, capabilities); server.add_tool("hello", "Greets the user", {{"type", "object"}, {"properties", {{"name", {{"type", "string"}}}}}}, - [](const nlohmann::json& args) { - return {{"message", "Hello, " + args["name"].get() + "!"}}; + [](const nlohmann::json& args) -> nlohmann::json { + return nlohmann::json{ + {"message", "Hello, " + args["name"].get() + "!"}}; }); - server.run_stdio(); // Blocks until connection closes + server.run_http("127.0.0.1", 3000); // Serves http://127.0.0.1:3000/mcp } ``` @@ -57,47 +70,100 @@ int main() { ```cpp #include -#include +#include -boost::asio::co_spawn(executor, [&]() -> mcp::Task { - auto transport = std::make_unique(executor); - mcp::Client client(std::move(transport), executor); +#include +#include +#include +#include +#include - co_await client.connect({"my-client", "1.0.0"}, {}); - auto result = co_await client.call_tool("hello", {{"name", "World"}}); - std::cout << result.content.dump() << std::endl; -}, boost::asio::detached); +int main() { + boost::asio::io_context io; + auto transport = std::make_shared( + io.get_executor(), "http://127.0.0.1:3000/mcp"); + mcp::Client client(transport, io.get_executor()); + + boost::asio::co_spawn(io, [&]() -> mcp::Task { + co_await client.connect("my-client", "1.0.0"); + nlohmann::json arguments{{"name", "World"}}; + auto result = co_await client.call_tool("hello", arguments); + std::cerr << nlohmann::json(result).dump(2) << '\n'; + client.close(); + }, boost::asio::detached); + + io.run(); +} ``` +Note the named `arguments` variable. See [Compiler +Notes](#compiler-notes) for why it is not built inline. + ## Usage Highlights ### Tools, Resources, and Prompts -```cpp -// Register a read-only resource -server.add_resource("mcp://status", "System status", "text/plain", []() { - return "All systems go."; -}); - -// Register a prompt template -server.add_prompt("greet", "Greets the user", {{"name", "User name"}}, [](const nlohmann::json& args) { - return {{"messages", {{{"role", "user"}, {"content", {{"type", "text"}, {"text", "Hello, " + args["name"].get()}}}}}}}; -}); -``` +Resources and prompts use the same typed-handler model as tools: provide the +protocol metadata (`mcp::Resource`, `mcp::ResourceTemplate`, or `mcp::Prompt`) +and a handler whose input and output are serializable protocol types. See the +[stdio server example](examples/servers/stdio/server_stdio.cpp) for complete, +compiled registrations. ### Server Context (Logging & Progress) Async handlers have access to a `Context` for real-time interaction: ```cpp -server.add_tool("long_task", "A task with progress", schema, - [](const nlohmann::json& args, mcp::Context& ctx) -> mcp::Task { - ctx.log_info("Starting work..."); +server.add_tool("long_task", "A task with progress", schema, + [](mcp::Context& ctx, const nlohmann::json& args) -> mcp::Task { + co_await ctx.log_info("Starting work..."); co_await ctx.report_progress(50, 100); - co_return {{"status", "done"}}; + co_return nlohmann::json{{"status", "done"}}; }); ``` +## Compiler Notes + +### GCC 12 and 13: initializer lists inside a `co_await` expression + +GCC 12 and GCC 13 crash with an internal compiler error when a `co_await` +expression contains an initializer list whose elements have non-trivial +destructors. GCC 13 is the default compiler on Ubuntu 24.04, so this is easy to +hit on a stock toolchain. The error looks like this, and names the closing brace +of the enclosing lambda rather than the offending argument: + +``` +internal compiler error: in build_special_member_call, at cp/call.cc:11096 +``` + +Build the container into a named variable first: + +```cpp +// Crashes GCC 12 and GCC 13 +auto result = co_await client.call_tool("hello", nlohmann::json{{"name", "World"}}); +auto prompt = co_await client.get_prompt("greet", std::map{{"who", "you"}}); + +// Compiles +nlohmann::json arguments{{"name", "World"}}; +auto result = co_await client.call_tool("hello", arguments); +``` + +The trigger is the initializer list, not the type, so it is not specific to +`nlohmann::json`: `std::map{{"a", "b"}}` and +`std::vector{"a", "b"}` crash the same way, while +`std::vector{1, 2, 3}` and `nlohmann::json::object()` do not. Parentheses +do not help, because `nlohmann::json({{"a", 1}})` still forms an initializer +list. Only hoisting the value out of the `co_await` expression avoids it. + +This is a compiler defect rather than an SDK one, and no change to the SDK's +signatures avoids it: taking the argument by value instead of by reference +still crashes. + +The contributing guide's [Known +Issues](https://yurirocha15.github.io/mcp-cpp-sdk/contributing.html#known-issues) +section carries the full case list, alongside the separate GCC 11 coroutine bug +this codebase also works around. + ## Documentation For full guides, API reference, and integration details, visit our **[Documentation Site](https://yurirocha15.github.io/mcp-cpp-sdk)**. @@ -132,6 +198,7 @@ python scripts/build.py --examples --test | `--debug` | Build in debug mode | | `--test` | Build and run unit tests | | `--examples` | Build example applications | +| `--conformance` | Build fixtures for the official MCP conformance runner | | `--linkage {both,shared,static}` | Select which SDK linkage variants to build | | `--cppstd {20,23}` | Select the C++ consumer standard (default: C++20) | | `--sanitize` | Build with ASan/UBSan (Linux/macOS) | @@ -148,6 +215,11 @@ python scripts/build.py --examples --test Please see the [CONTRIBUTING guide](docs/contributing.rst) for the full process. +Project maintenance commitments and release gates are documented in the +[maintenance policy](MAINTENANCE.md), [dependency policy](DEPENDENCY_POLICY.md), +[versioning policy](VERSIONING.md), [Tier roadmap](ROADMAP.md), and +[security policy](SECURITY.md). + ## License Apache License 2.0 - see [LICENSE](LICENSE) for details. diff --git a/ROADMAP.md b/ROADMAP.md new file mode 100644 index 0000000..240f001 --- /dev/null +++ b/ROADMAP.md @@ -0,0 +1,53 @@ +# Roadmap to MCP SDK Tier 1 + +This repository is an independent community SDK. Completing these gates makes +the project tier-ready; official Tier assignment and inclusion in the MCP SDK +roster require approval from MCP governance. + +## Tier 2 foundation + +- Pin the official conformance runner and publish separate client and server + results for protocol revision `2025-11-25`. +- Reach at least 80% applicable conformance on both sides, with 100% as this + project's internal target and all expected failures kept visible. +- Close protocol lifecycle, capability negotiation, structured result, + timeout/cancellation, disconnect cleanup, and HTTP security gaps. +- Compile documentation examples and validate installed-package consumers in + CI. +- Operate the 30-day issue-triage and 14-day P0 maintenance commitments. +- Implement newly released non-experimental protocol features within six + months, with their conformance, documentation, and example coverage. +- Publish a non-prerelease `1.0.0` or later. Project-stable `0.x` releases do + not satisfy the official Tier 2 stable-release requirement. +- Complete the public API and dependency review required for `1.0.0`. + +## Stable 1.0 release + +- Freeze the supported C++ API, compatibility policy, compiler matrix, and + dependency floors. +- Publish a non-prerelease `1.0.0` with migration notes and immutable + conformance evidence. +- Rebaseline against the final 2026 protocol and the conformance release chosen + by the MCP SDK Working Group before making a tier application. + +## Tier 1 readiness + +- Maintain 100% of applicable client and server conformance for the accepted + current protocol revision. +- Provide runnable documentation and examples for all 48 non-experimental MCP + features in the Tier assessment, with a checked coverage matrix linking each + feature to its API, guide, and example. +- Demonstrate issue triage within two business days and P0 resolution within + seven calendar days. +- Track MCP release candidates early enough to support required features on the + release schedule agreed with the SDK Working Group. +- Maintain explicit source, ABI, deprecation, dependency, and breaking-change + policies for every stable release. + +## Governance checkpoint + +Before claiming an official tier, maintainers will ask the MCP SDK Working +Group to clarify roster admission, repository governance, the pinned +conformance release, and evidence submission. Until accepted, project material +will report measured conformance percentages or use "tier-ready" language, +never an official Tier 1 or Tier 2 designation. diff --git a/VERSIONING.md b/VERSIONING.md new file mode 100644 index 0000000..6eb308f --- /dev/null +++ b/VERSIONING.md @@ -0,0 +1,52 @@ +# Versioning and Compatibility Policy + +`mcp-cpp-sdk` follows Semantic Versioning. The `VERSION` file is authoritative, +and release tags use `vMAJOR.MINOR.PATCH` or `vMAJOR.MINOR.PATCH-rc.NUMBER`. +Protocol revision numbers and SDK release numbers are independent. + +## Public compatibility surface + +The compatibility surface includes installed headers, exported CMake and +pkg-config targets, documented compiler and dependency floors, serialized MCP +behavior, and documented command-line interfaces. Test helpers, benchmark +adapters, conformance fixtures, source files under `src/`, and symbols in a +`detail` namespace are not public API. + +## Before 1.0 + +- Minor releases may make breaking source or behavior changes when the release + notes include migration guidance. +- Patch releases contain compatible fixes and documentation changes. +- Release candidates are not stable and may change before the matching final + release. + +The project will not publish `1.0.0` until client and server conformance, +public-API review, packaging, and the current protocol transition meet the +gates in `ROADMAP.md`. + +## From 1.0 onward + +- Major releases may contain breaking changes. +- Minor releases add functionality while preserving supported source and wire + compatibility. +- Patch releases contain compatible fixes and security updates. +- The shared library's major SOVERSION identifies its ABI compatibility line. + Consumers that require ABI stability should use a matching major line and a + supported toolchain configuration. + +Public APIs are deprecated in headers and release notes before removal. Except +for urgent security or protocol-correctness fixes, removal occurs no sooner +than the next major release and after at least one minor release or 90 days, +whichever is longer. Protocol features follow the MCP feature lifecycle; +deprecated features remain available for their required compatibility window +and include migration guidance. + +## Release evidence + +The release gate requires an updated `CHANGELOG.md`, passing platform and +package-consumer CI, recorded supported MCP protocol revisions, and reviewed +client/server conformance results. The current CI artifacts are regression +evidence with limited retention, not an immutable release archive. Linking +durable conformance evidence from every release becomes mandatory once release +automation archives it. Any intentional compatibility exception is called out +explicitly in the release notes. diff --git a/benchmark/.dockerignore b/benchmark/.dockerignore new file mode 100644 index 0000000..12e107f --- /dev/null +++ b/benchmark/.dockerignore @@ -0,0 +1,4 @@ +results/ +graphify-out/ +benchmark-mcp-servers-v2/ +alternative-sdks/*/.git/ diff --git a/benchmark/.gitignore b/benchmark/.gitignore new file mode 100644 index 0000000..51a5b9c --- /dev/null +++ b/benchmark/.gitignore @@ -0,0 +1,3 @@ +/benchmark-mcp-servers-v2/ +/alternative-sdks/ +/results/*/ diff --git a/benchmark/README.md b/benchmark/README.md index 5f15460..1601be4 100644 --- a/benchmark/README.md +++ b/benchmark/README.md @@ -1,235 +1,369 @@ -# MCP Benchmark — TM Dev Lab v2 +# MCP benchmark -Performance comparison of C++, Python, Go and Rust MCP server implementations under identical I/O-bound workloads (Redis + HTTP). +This directory contains a reproducible, end-to-end benchmark for MCP servers. +The workload sends MCP requests over HTTP and exercises the same Redis and HTTP +application operations for every comparable C++ SDK. -Methodology mirrors [TM Dev Lab v2](https://github.com/thiagomendes/benchmark-mcp-servers-v2). -The Python, Go, and Rust servers — as well as the API service, Redis seeder, and k6 script — are sourced directly from that upstream repo (pinned to commit `8a9a5f8e`). - ---- +The application workload is based on +[TM Dev Lab v2](https://github.com/thiagomendes/benchmark-mcp-servers-v2), pinned +to commit `8a9a5f8ef505f46b6079072ef4603304ca672e33`. That repository supplies the API +service, Redis seeder, and language-baseline server sources. Load generation +uses the audited local profile in `alternatives/benchmark.js`; it does not use +the upstream k6 script. ## Prerequisites -- [Docker](https://docs.docker.com/get-docker/) with Compose v2 (`docker compose`) -- `python3`, `jq`, and `git` (for orchestration, results parsing, and cloning upstream) - -> k6 runs inside a Docker container (`grafana/k6`) — no host installation needed. +- A Linux host with Docker cgroup accounting enabled; the harness reads Linux + CPU affinity and cgroup files and uses host networking +- Docker with Compose v2 (`docker compose`) +- `python3`, `jq`, and `git` ---- +k6 runs from a container image pinned by digest, so a host k6 installation is +not required. A full-duration run also requires at least nine CPUs in the +orchestrator's effective affinity set; this leaves measurable capacity for the +8.5 CPUs assigned across the target and shared services. -## Quick Start +## Quick start ```bash -cd benchmark/ +cd benchmark -# Benchmark all four servers (builds, seeds Redis, warms up, runs k6) -./run.sh +# Run the five comparable C++ SDKs. +./run.sh cpp-sdks -# Benchmark specific servers -./run.sh cpp -./run.sh cpp python -./run.sh cpp,go -``` +# Run the language baselines as a diagnostic (not a publishable ranking). +./run.sh baseline -`run.sh` handles everything end-to-end: -1. Clones the upstream benchmark repo (once, pinned to commit `8a9a5f8e`) into `benchmark/benchmark-mcp-servers-v2/` -2. Starts Redis + API service -3. Seeds Redis with 130k keys (carts, history, popularity, rate limits) -4. For each selected server: resets Redis, starts only that server, warms up, runs k6 **3 times**, picks the median run -5. Collects Docker CPU/memory/network stats during the test -6. Prints a comparison table and saves results to `benchmark/results//` - -Results per server: -- `/k6_summary.json` — canonical (median) k6 metrics -- `/k6_summary_run{1,2,3}.json` — raw results from each of the 3 k6 runs -- `/k6_multi_run_stats.json` — per-run RPS and coefficient of variation % -- `/k6_console_run{1,2,3}.log` — k6 terminal output per run -- `/stats.json` — CPU/memory/network samples during the test -- `comparison.txt` — side-by-side RPS, latency percentiles, error rates - -### Benchmark Profile (TM Dev Lab v2) - -- **50 virtual users**, 5-minute sustained load -- 15s ramp-up, 10s ramp-down -- 60s warmup excluded from metrics (5 init sessions + 9 full tool sessions per server) -- Each VU cycles through all three tools + `tools/list` -- Redis FLUSHDB + re-seed between servers - ---- - -## Architecture - -```mermaid -graph TD - subgraph Docker network - Redis["Redis :6379"] - API["API Service :8100
(Go stdlib, 100k products)"] - CPP["C++ MCP :8080"] - Python["Python MCP :8081"] - Go["Go MCP :8082"] - Rust["Rust MCP :8083"] - end - - Redis --- CPP - Redis --- Python - Redis --- Go - Redis --- Rust - API --- CPP - API --- Python - API --- Go - API --- Rust +# Run selected servers. +./run.sh ours-comparable hkr04 ``` -Each MCP server exposes the same three tools: +The publishable C++ comparison uses the production defaults: three measured +runs per SDK, 50 virtual users, and five minutes of constant load per measured +run. Environment variables exposed by `run.sh` may shorten a local smoke test, +but results from a shortened profile must not be published as benchmark data. +Only one `run.sh` invocation may use the host at a time. An exclusive lock in +`/tmp` rejects overlapping runs across worktrees before they can share the +fixed containers or ports. -| Tool | Operations | -|---|---| -| `search_products` | Parallel: HTTP product search + Redis `ZREVRANGE` (popularity) | -| `get_user_cart` | Sequential Redis `HGETALL` (cart), then parallel: HTTP product lookup + Redis `LRANGE` (history) | -| `checkout` | Parallel: HTTP cart total + Redis `INCR` (rate limit), then sequential `RPUSH` + `ZINCRBY` | +The result-directory suffix and `run_manifest.json` identify the effective +profile: ---- - -## Server Ports - -| Service | Port | Notes | +| Invocation | Result profile | Publication status | |---|---|---| -| Redis | 6379 | Internal | -| API service | 8100 | Go stdlib, no external deps | -| C++ MCP | 8080 | MCP + `/health` on same port (Streamable HTTP) | -| Python MCP | 8081 | MCP + `/health` on same port | -| Go MCP | 8082 | MCP + `/health` on same port | -| Rust MCP | 8083 | MCP + `/health` on same port | - -> All four servers expose MCP and health endpoints on the same port. - ---- - -## Correctness Verification - -`run.sh` uses the pinned upstream k6 script for both load generation and response -validation. Each k6 run checks `initialize`, `tools/list`, and all benchmark -tool responses before writing the retained summary files. - ---- - -## Manual Operations - -> **Note:** `docker-compose.yml` builds the Python/Go/Rust servers, API service, and Redis seeder from the upstream clone at `benchmark/benchmark-mcp-servers-v2/`. Run `./run.sh` once first to ensure the clone exists, or clone manually: -> ```bash -> git clone https://github.com/thiagomendes/benchmark-mcp-servers-v2.git benchmark/benchmark-mcp-servers-v2 -> cd benchmark/benchmark-mcp-servers-v2 && git checkout 8a9a5f8ef505f46b6079072ef4603304ca672e33 -> ``` - -### Start individual servers - -Redis and API service are always required: - -```bash -cd benchmark/ - -# Just Redis + API + C++ -docker compose up redis api-service cpp-server - -# Just Redis + API + Python -docker compose up redis api-service python-server - -# Just Redis + API + Go -docker compose up redis api-service go-server - -# Just Redis + API + Rust -docker compose up redis api-service rust-server -``` - -### Seed Redis manually - -```bash -docker compose --profile seeder up redis-seeder -``` - -### Manual MCP tool call +| `./run.sh cpp-sdks` with exact defaults | `production` | Publishable candidate after every gate passes | +| `./run.sh baseline` with exact defaults | `baseline-diagnostic` | Diagnostic only; never an equal-work ranking | +| Any parameter override or other server selection | `smoke` | Harness validation only | + +`publishable_candidate` remains false while a run is in progress and after any +failure. It becomes true only when the exact production profile reaches +`status: "complete"`; consumers must require both fields. Generated result +directories are ignored by default so retaining or publishing an audited bundle +requires an explicit action. + +## Comparable C++ scope + +`cpp-sdks` contains five public-API adapters: + +| Benchmark name | SDK | Host port | +|---|---|---:| +| `ours-comparable` | this SDK | 8089 | +| `hkr04` | hkr04/cpp-mcp | 8084 | +| `fastmcpp` | FastMCPP | 8085 | +| `cxxmcp` | cxxmcp | 8086 | +| `neumann` | Neumann-Labs/mcp-cpp | 8087 | + +External source repositories and exact commits are recorded in +`alternatives/sources.tsv`. The adapters use each SDK's public API. They do not +copy, vendor, patch, or bypass another SDK's protocol or server implementation. + +The harness pins direct source commits and top-level container images, but it +does not claim a bit-for-bit hermetic rebuild. Ubuntu package repositories and +transitive dependencies resolved by the alternatives' CMake builds are not all +content-addressed; their resolved image identities are therefore captured with +each run. + +Gopher MCP remains recorded in the source manifest but excluded from this +group. Its public server transport uses legacy HTTP+SSE rather than the single +Streamable HTTP endpoint exercised here, so including it would test a different +transport contract. cxxmcp is included: the local client sends the required +negotiated protocol-version header after initialization. + +The baseline group is separate from the comparable C++ group. The pinned +upstream repository supplies the Python, Go, and Rust baseline servers and the +shared infrastructure. The optimized `cpp` baseline is retained as a separate +application implementation and must not be presented as an SDK-adapter result. + +Python, Go, and Rust are rerun after harness changes, but their results remain +diagnostic. Their upstream Docker builds contain floating inputs: Python uses +version ranges, Go regenerates dependency resolution during the image build, +Rust does not copy its lockfile into the build, and the baseline Dockerfiles use +unpinned base-image tags. Exact built image identities are retained with every +run, but these inputs prevent a source commit from defining a bit-for-bit +rebuild. + +Every target uses the same `upstream-v2-strict-mcp-v1` measurement eligibility +contract. It preserves the pinned benchmark's observable business predicates +while adding uniform MCP lifecycle, zero-error, resource, and Redis-side-effect +gates. The five C++ adapters also must pass `adapter-exact-v1` during preflight +and postflight because their glue and shared workload are controlled here; that +supplemental check does not execute in the measured k6 path. + +The supplemental diagnostics intentionally expose two upstream limitations +without excluding otherwise working language baselines. Python accepts the +canonical checkout call but advertises unconstrained array elements for its +untyped `items: list` parameter. Rust executes all three Redis mutations in one +pipeline, but returns the pipeline's `ZADD` result as `rate_limit_count` instead +of the `INCR` result. Both pass the inherited observable predicates and the +direct Redis side-effect gate; Python and Rust simply record the corresponding +`adapter-exact-v1` diagnostic as false. These interface, build, and execution +differences are why language baselines must not appear in the publishable C++ +SDK ranking. + +## Workload + +Every comparable adapter registers the same three tools and uses the shared +`benchmark_workload` implementation: + +| Tool | End-to-end operations | +|---|---| +| `search_products` | HTTP product search in parallel with Redis popularity lookup | +| `get_user_cart` | Redis cart lookup followed by HTTP product lookup and Redis history lookup | +| `checkout` | HTTP total calculation, Redis rate-limit increment, history append, and popularity update submitted concurrently | + +Each virtual user repeatedly opens MCP sessions, calls all three tools, calls +`tools/list`, and closes stateful sessions. The headline operations rate counts +`tools/call` and `tools/list` responses only after their shape and universal +contract checks have run. Raw HTTP request rate is reported separately because it also +includes initialization, notification, and session cleanup traffic. Tool-call +latency is collected in one combined k6 trend across the three tools, with each +tool also reported separately. + +This is an end-to-end Redis/HTTP workload. It is not a parser-only, +serialization-only, or in-process SDK microbenchmark. + +## Concurrency and resource controls + +Each comparable server container is limited to 2 CPUs and 2 GiB of memory. +Top-level blocking request/handler capacity is normalized to 50, matching the +production profile's 50 VUs. This prevents an HTTP implementation whose workers +own persistent keep-alive connections from being limited to fewer active +clients than the load profile. It is an admission-capacity control, not a claim +that every process has the same number of OS threads or 50 CPUs. + +All five adapters call the same application workload, whose internal pool has +64 threads for the Redis and HTTP work. This SDK retains two asynchronous I/O +threads. hkr04's separate two-thread async pool is not used by its synchronous +Streamable HTTP path. SDK/runtime-owned I/O, session, timer, logging, and other +auxiliary threads otherwise remain native to each implementation and can +differ. Those differences are not hidden; they are part of the process measured +by Docker CPU, memory, and network sampling. The harness records the effective +container limits and image identity for every measured run. + +Redis is limited to 0.5 CPU and 512 MiB, the API service to 2 CPUs and 2 GiB, +and k6 to 4 CPUs and 2 GiB. Every measured container is constrained to the same +host CPU-affinity set inherited by the harness. The preflight rejects a Docker +or effective cgroup CPU set that differs from that recorded set. + +One-second collectors cover the target, Redis, API service, k6, and host for +the measurement window. Collector audit files must show complete sampling, +zero failed attempts, sufficient duration, and bounded gaps. The target's +utilization is reported rather than rejected because target saturation is a +result; Redis, API, k6, and host p95 CPU and memory must each remain below 90% +of the recorded limit or the run is invalid. + +Host network counters use the interface cohort present when measurement starts. +If Docker retires one of those interfaces during the run, its last counters are +frozen so the aggregate stays monotonic and the retirement is recorded in the +collector audit. Interfaces created after the baseline are not adopted, and a +counter reset on an interface that remains present still invalidates collection. + +## Corrected run procedure + +The C++ SDK comparison follows this sequence: + +1. Pin and verify every source checkout, build the selected images once, and + capture host, toolchain, source-tree, image, and Compose provenance. +2. Start the shared Redis and API services and generate a deterministic, + seed-recorded counterbalanced order for three rounds. Only one MCP server is + under load at a time. +3. Before each SDK/run pair, reset and re-seed Redis and run the standalone + protocol and workload verifier. It checks the Redis counter, history, and + popularity deltas around checkout; C++ adapters also run supplemental exact + adapter validation. +4. Run a separate warmup invocation: ramp from 0 to 50 VUs for 15 seconds, then + hold 50 VUs for 60 seconds. Warmup metrics are never merged into measured + metrics. Ramp-down and graceful-stop time are both zero. +5. Reset and re-seed Redis again so the measurement starts from the canonical + fixture state. +6. Start per-run resource collection and run exactly five minutes at a constant + 50 VUs. There is no measured ramp-up or ramp-down phase, and graceful-stop + time is zero, so the load generator cannot extend the measurement boundary. +7. Stop resource collection immediately after k6 exits, re-seed, and verify the + server again. Any eligibility, required supplemental, or correctness-threshold + failure makes that run fail. +8. After all three counterbalanced rounds, select each SDK's median run by + measured operations per second. Copy the resource samples and container/image + evidence from that same run; do not combine resources from another run. + +The counterbalanced schedule varies server position and neighboring servers +across rounds, reducing fixed-order thermal and cache bias. Its seed and exact +per-round order are retained with the result bundle. + +Sessions unfinished at the hard time boundary are right-censored. An operation +is counted only after its HTTP response, shape check, contract check, and metric +recording complete; a response interrupted before that point is not counted. +An unfinished session contributes neither a completion result nor full-session +latency. Check totals may exceed the conservative operation counter if k6 stops +a VU in the tiny interval after its checks but before the counter update. +Started and completed session counts are reported separately so this boundary +behavior remains visible. + +## Protocol and correctness gate + +The client requests MCP revision `2024-11-05` during `initialize`. Each server +may select a supported revision; that selected revision is recorded in the +preflight, postflight, k6 summary, median-run metadata, and comparison table. +It must remain stable across all rounds for that server, and every subsequent +request uses the selected value in `MCP-Protocol-Version`. + +The standalone verifier and local k6 profile use the same universal eligibility +contract for every target. They check: + +- the negotiated `initialize` result; +- acceptance of `notifications/initialized`; +- exactly the three expected tools from `tools/list`, each with an input-schema + object, followed by successful canonical calls; +- the pinned workload's observable predicates: search count/list sizes, the + requested user's nonempty cart and five history entries, and a confirmed + two-item checkout with a positive total and numeric rate-limit field; +- direct pre/post checkout Redis evidence that the rate counter, history length, + and product popularity score each increase by one; and +- session deletion when the server issued a session ID. + +The gate also validates JSON-RPC response identity, MCP response media type and +framing, session headers, and the required empty `202` response to the +initialized notification. Post-initialization requests include +`MCP-Protocol-Version` with the selected version and `Mcp-Session-Id` when +applicable. The k6 thresholds require a 100% check pass rate, zero HTTP request +failures, and zero MCP errors. Threshold evaluation aborts a load phase as soon +as a violation is observed because that phase is already invalid; it is never +eligible for a comparison table. + +`adapter-exact-v1` is separate, out-of-band validation. It checks detailed input +schemas, deterministic rows and histories, exact checkout fields, and +server-type markers. Its outcome is recorded for every target and required for +the five authored C++ adapters. It is advisory for the pinned language +baselines, so schema precision or a mislabeled response field cannot silently +change the inherited workload eligibility. Because this harness counts only +completed `tools/call` and `tools/list` operations, its operations/s values must +not be compared directly with the upstream repository's RPS metric, which also +counted initialization requests. + +## Reported metrics + +- Operations per second: `tools/call` plus `tools/list`, excluding lifecycle + messages from the headline rate +- Raw HTTP requests per second, including lifecycle traffic +- Combined tool-call p50/p90/p95/p99 and per-tool latency distributions +- Session success/failure counts and strict correctness/error rates +- Per-run Docker CPU, memory, and network samples +- Three-run sample coefficient of variation and, for the default three runs, a + 95% t-interval for mean operations rate + +Percentiles come directly from the selected run's raw k6 trend. Percentiles are +not averaged or reconstructed from pre-aggregated per-tool quantiles. + +## Result and provenance files + +Results are written under `results/_/`; a baseline run uses +the explicit `baseline-diagnostic` suffix. Before any checkout or image build, +the harness stages read-only copies of the resolved runtime inputs under +`harness/`: `docker-compose.yml`, `benchmark.js`, and `SHA256SUMS`. Their +in-memory expected hashes are checked before Compose and k6 operations and once +more before completion. + +The root bundle also contains `source_snapshot.tar`, `environment.json`, +`run_order.json`, `compose.resolved.yml`, `compose_images.json`, `build.log`, and +`run_manifest.json`. The manifest names the universal eligibility contract and +the targets for which supplemental adapter validation is required. The +deterministic source snapshot covers the local SDK, +adapters, and harness inputs; its digest is checked after builds and again +before aggregation and completion. Environment metadata records the local +worktree state, exact alternative and upstream commits/origins, host affinity, +tool versions, and pinned k6 image identity. Each server directory retains, for +every round: + +- `warmup_summary_runN.json` and `warmup_console_runN.log`; +- `k6_summary_runN.json` and `k6_console_runN.log`; +- `protocol_preflight_runN.json` and `protocol_postflight_runN.json`, with + eligibility checks, supplemental diagnostics, Redis side-effect observations, + and matching `.log` files that preserve verifier output even when JSON + evidence cannot be produced; +- `stats_runN.json`, `redis_stats_runN.json`, `api_stats_runN.json`, + `k6_stats_runN.json`, and `host_stats_runN.json`, each paired with an + `.audit.json` collector record. Host audits identify the initial, active, + retired, and ignored-new interface sets under the recorded + `baseline-cohort-retire-v1` policy; +- `resource_headroom_runN.json` with the integrity and headroom verdict; +- `container_inspect_runN.json`, `image_inspect_runN.json`, + `cgroup_runN.json`, and `server_runN.log`; and +- matching k6 container, image, and cgroup evidence under `k6/`. + +After median selection, `k6_summary.json`, `stats.json`, +`container_inspect.json`, and `image_inspect.json` are copies from the same +selected run. The selected protocol, cgroup, resource-headroom, collector-audit, +and k6 provenance records are paired the same way. `resource_summary.json` +summarizes only those selected resource samples. Its host network observation +is explicitly scoped to the baseline interface cohort and labels coverage +partial when an interface retires or a new interface is ignored. +`k6_multi_run_stats.json` records the three measured rates and selection +statistics. +Failed manifests also record the failing stage, exit code, final diagnostic, +and relative evidence-log path when the verifier identified the failure. + +## Manual protocol check + +The benchmark requests protocol revision `2024-11-05`. An initialize request +does not need a protocol-version header: ```bash -# 1. Initialize session (C++ server) -curl -s -X POST http://localhost:8080/mcp \ - -H "Content-Type: application/json" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ +curl -i -sS -X POST http://localhost:8089/mcp \ + -H 'Content-Type: application/json' \ + -H 'Accept: application/json, text/event-stream' \ -d '{ - "jsonrpc": "2.0", - "id": 1, - "method": "initialize", - "params": { - "protocolVersion": "2025-11-25", - "clientInfo": {"name": "test", "version": "1.0"}, - "capabilities": {} - } - }' - -# 2. Call search_products (use MCP-Session-Id from step 1 response headers) -curl -s -X POST http://localhost:8080/mcp \ - -H "Content-Type: application/json" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ - -H "MCP-Session-Id: " \ - -d '{ - "jsonrpc": "2.0", - "id": 2, - "method": "tools/call", - "params": { - "name": "search_products", - "arguments": {"category": "Electronics", "min_price": 50, "max_price": 500, "limit": 5} + "jsonrpc":"2.0", + "id":1, + "method":"initialize", + "params":{ + "protocolVersion":"2024-11-05", + "clientInfo":{"name":"manual-check","version":"1.0"}, + "capabilities":{} } }' ``` ---- - -## Teardown +Read the server-selected version from the initialize result. Use that value and +the returned `Mcp-Session-Id` on every later request: ```bash -# Stop and remove containers (keeps Redis data if volume was configured) -docker compose down - -# Full cleanup including images -docker compose down --rmi all --volumes +curl -sS -X POST http://localhost:8089/mcp \ + -H 'Content-Type: application/json' \ + -H 'Accept: application/json, text/event-stream' \ + -H 'MCP-Protocol-Version: ' \ + -H 'Mcp-Session-Id: ' \ + -d '{"jsonrpc":"2.0","method":"notifications/initialized"}' ``` ---- - -## Troubleshooting - -**Upstream clone is missing** — if you see build errors like `unable to prepare context: path not found`, the upstream repo hasn't been cloned yet. Run `./run.sh` once to auto-clone it, or clone manually (see [Manual Operations](#manual-operations)). - -**C++ server build is slow** — the first build compiles the full SDK via Conan inside Docker. Subsequent builds use the Docker layer cache. Expect 3–5 minutes on first run. +## Teardown -**`healthy` never appears for cpp-server** — the healthcheck pings port 8080. If it isn't reachable, the server process likely failed during startup. Check logs: ```bash -docker compose logs cpp-server -``` +# Stop and remove benchmark containers. +docker compose down -**Redis seeder exited with error** — ensure Redis is healthy before the seeder runs. If re-seeding, flush Redis first: -```bash -docker compose exec redis redis-cli FLUSHALL -docker compose --profile seeder up redis-seeder +# Optional full cleanup, including images and volumes. +docker compose down --rmi all --volumes ``` -**Port conflicts** — if 8080/8081/8082/8100/6379 are in use locally, edit the host-side port mappings in `docker-compose.yml` (left side of `host:container`). - ---- - -## Results - -See [RESULTS.md](RESULTS.md) for the retained benchmark records: - -- `20260405_205033`: C++, Go, and Python comparison -- `20260620_220910`: C++ three-run verification - -The `20260620_220910` C++ benchmark was run three times by `run.sh cpp`; the median run achieved **7,025.13 RPS** with **0.22% CV**, **0% errors**, **6.92 MB average memory**, and **8.15 MB max memory**. - -### Fair Comparison Status - -The `20260405_205033` results compare the retained C++, Go, and Python artifacts. The `20260620_220910` results verify the C++ implementation across three runs. - -1. **Identical Infrastructure**: All servers use the same upstream API service, Redis seeder, and Docker resource limits. -2. **Methodology Parity**: The k6 benchmark script matches upstream methodology exactly. -3. **Hardware Consistency**: All tests run on the same hardware (AMD Ryzen 9 9900X). +See [RESULTS.md](RESULTS.md) for publication status. A reduced smoke test proves +that the harness works; it is not a performance result. diff --git a/benchmark/RESULTS.md b/benchmark/RESULTS.md index fa00bb8..25038b8 100644 --- a/benchmark/RESULTS.md +++ b/benchmark/RESULTS.md @@ -1,53 +1,86 @@ -# Benchmark Results - -This file keeps the benchmark records by run date: - -- `20260405_205033`: C++, Go, and Python comparison -- `20260620_220910`: C++ three-run verification - -## Test Profile - -| | | -|---|---| -| **Workload** | TM Dev Lab v2 MCP benchmark: Redis + HTTP I/O-bound tools | -| **Load** | 50 VUs, 15s ramp-up, 5m sustained load, 10s ramp-down | -| **Warmup** | 60s warmup excluded from metrics | -| **Repetition** | `run.sh` executes 3 k6 runs and selects the median by RPS | -| **Infrastructure** | Same upstream API service and Redis seeder, Docker Compose | -| **Host** | AMD Ryzen 9 9900X, 32 GB RAM, Ubuntu kernel 6.17.0-20-generic | - -## 20260405_205033 - -Results from `benchmark/results/20260405_205033/`. - -| Server | RPS | p50 (ms) | p99 (ms) | Error Rate | Avg Memory | Max Memory | -|---|---:|---:|---:|---:|---:|---:| -| C++ | 12,191.61 | 0.31 | 5.55 | 0% | 11.34 MB | 13.07 MB | -| Go | 9,154.28 | 0.36 | 36.06 | 0% | 21.38 MB | 24.52 MB | -| Python | 904.33 | 18.17 | 190.20 | 0% | 58.09 MB | 61.86 MB | - -## 20260620_220910 - -C++-only verification from `benchmark/results/20260620_220910/`. - -| Run | RPS | -|---:|---:| -| 1 | 6,995.45 | -| 2 | 7,025.13 | -| 3 | 7,031.65 | - -| Metric | Value | -|---|---:| -| Median run | 2 | -| Median RPS | 7,025.13 | -| Mean RPS | 7,017.41 | -| CV | 0.22% | -| Requests | 2,283,376 | -| p50 | 0.68 ms | -| p95 | 1.92 ms | -| p99 | 2.69 ms | -| Error rate | 0% | -| Avg memory | 6.92 MB | -| Max memory | 8.15 MB | - -Compared with the `20260405_205033` C++ result, the `20260620_220910` median C++ run used less Docker-reported memory: average memory decreased from **11.34 MB** to **6.92 MB**, and max memory decreased from **13.07 MB** to **8.15 MB**. +# Benchmark results + +## Publication status + +The five-C++ production run `20260720_012851_production` completed every gate +and passed an independent artifact audit. Its manifest is `complete` with +`publishable_candidate: true`. + +The evidence bundle is retained locally at +`benchmark/results/20260720_012851_production/`. Generated result directories +are intentionally ignored by Git; an external citation should include an +unchanged copy of that bundle. The SHA-256 of its sorted per-file checksum +stream is +`e0d1807913cba76053dc9ba1596ac5474b29ac33fa653e4f39a9eb9e3ba72316`. + +## Production result + +The run used three counterbalanced rounds (order seed 417), 50 VUs, a separate +15-second ramp plus 60-second warmup, and an exact five-minute constant-load +measurement for every SDK/round pair. Each target had the same 2-CPU, 2-GiB, +and 50-request admission limits. Every row used the shared +`upstream-v2-strict-mcp-v1` measured contract; required exact C++ adapter checks +ran only in pre/postflight. + +| C++ SDK | Protocol | Median ops/s | Sample CV | p50 ms | p95 ms | p99 ms | Target p95 CPU | Target p95 memory | Errors | +|---|---:|---:|---:|---:|---:|---:|---:|---:|---:| +| This SDK | 2024-11-05 | 3,523.41 | 0.12% | 0.49 | 1.76 | 2.56 | 174.54% | 43.27 MiB | 0 | +| hkr04/cpp-mcp | 2025-03-26 | 678.33 | 0.04% | 40.77 | 41.83 | 41.93 | 22.12% | 32.65 MiB | 0 | +| FastMCPP | 2024-11-05 | 1,123.07 | 0.03% | 0.38 | 41.77 | 42.07 | 42.94% | 26.33 MiB | 0 | +| cxxmcp | 2024-11-05 | 2,987.82 | 0.56% | 1.64 | 6.12 | 7.67 | 72.46% | 30.14 MiB | 0 | +| Neumann-Labs/mcp-cpp | 2025-11-25 | 676.07 | 0.05% | 40.80 | 41.85 | 42.09 | 41.62% | 56.34 MiB | 0 | + +CPU follows Docker's convention: 100% is one fully used core, and each target +had a two-core limit. Resource values are from the same run selected for that +SDK's median throughput. Target saturation is reported, while headroom gates +apply to Redis, the API service, k6, and the host; every shared-resource gate +passed. + +The measured operation rates for all rounds were: + +| C++ SDK | Round 1 | Round 2 | Round 3 | Selected round | +|---|---:|---:|---:|---:| +| This SDK | 3,516.55 | 3,524.21 | 3,523.41 | 3 | +| hkr04/cpp-mcp | 678.49 | 678.33 | 677.99 | 2 | +| FastMCPP | 1,122.93 | 1,123.66 | 1,123.07 | 3 | +| cxxmcp | 2,987.82 | 2,959.55 | 2,988.71 | 1 | +| Neumann-Labs/mcp-cpp | 675.50 | 676.07 | 676.15 | 2 | + +These numbers describe this pinned MCP/Redis/HTTP workload and environment; +they are not a general-purpose application-performance claim. + +## Audit evidence + +The completed bundle records and passed: + +- 30/30 protocol gates (preflight and postflight for 15 measured rounds), + including the universal contract, required supplemental adapter validation, + stable negotiation, and direct Redis counter/history/popularity changes; +- 15/15 warmups and 15/15 measurements with zero MCP, HTTP, check, or session + failures; +- 75/75 collector audits and every per-round resource/headroom report; +- the `baseline-cohort-retire-v1` host network policy, with retirement and + ignored-new interface sets preserved and coverage labeled per selected run; + and +- final source, dependency, image, Compose, and immutable harness checks. + +The deterministic source tree hash is +`b8f31b64191f848271f8a12a4a42a352838f596c686de7d5493c61fecb967a6a`. +`source_snapshot.tar` hashes to +`7c87baf6498724fe47444c98c1b38be0871b3538734937af165ee22147a304dd`. +The archived k6 and Compose inputs hash to +`20393aa7219fac9473ff3ef601555a7619548d06ae59af258fe78ec144599b4e` +and `d100fc6d99d76d00a88eb0a7a0db2d1d43d4e589289d89a68f25e2af7add9292`. + +## Language diagnostic rerun + +The separate `20260720_030804_baseline-diagnostic` run also completed all nine +Rust, Go, and Python rounds with zero measured errors, valid universal +pre/postflight checks, direct Redis side effects, complete collectors, and +valid resource reports. It is intentionally not included in the production +table: its manifest is diagnostic-only and never an equal-work C++ ranking. + +The advisory layer preserves known upstream differences without changing +eligibility: Python emits a permissive array-item schema, and Rust returns a +mislabeled checkout count while independently producing the correct Redis +mutations. Neither difference was part of the inherited measured contract. diff --git a/benchmark/alternatives/CMakeLists.txt b/benchmark/alternatives/CMakeLists.txt new file mode 100644 index 0000000..343f92e --- /dev/null +++ b/benchmark/alternatives/CMakeLists.txt @@ -0,0 +1,161 @@ +cmake_minimum_required(VERSION 3.23) +project(mcp-cpp-sdk-alternative-benchmark LANGUAGES CXX) + +set(CMAKE_CXX_STANDARD 20) +set(CMAKE_CXX_STANDARD_REQUIRED ON) +set(CMAKE_CXX_EXTENSIONS OFF) + +# Match the production load's 50 VUs at the top-level request/handler boundary. +# cpp-httplib workers own keep-alive connections, so a smaller pool would limit +# active connections rather than normalize simultaneous request handling. Define +# both values because newer releases may grow the pool up to MAX_COUNT. +add_compile_definitions( + CPPHTTPLIB_THREAD_POOL_COUNT=50 CPPHTTPLIB_THREAD_POOL_MAX_COUNT=50 + MCP_BENCHMARK_REQUEST_CONCURRENCY=50) + +set(BENCHMARK_SDK + "" + CACHE STRING "SDK adapter id") +set(BENCHMARK_SDK_SOURCE + "" + CACHE PATH "Path to the SDK source") +if(NOT BENCHMARK_SDK OR NOT EXISTS "${BENCHMARK_SDK_SOURCE}/CMakeLists.txt") + message( + FATAL_ERROR "BENCHMARK_SDK and a valid BENCHMARK_SDK_SOURCE are required") +endif() + +find_package(CURL REQUIRED) +find_package(Threads REQUIRED) +find_library(HIREDIS_LIBRARY NAMES hiredis REQUIRED) +find_path(HIREDIS_INCLUDE_DIR NAMES hiredis/hiredis.h REQUIRED) + +add_library(benchmark-workload STATIC common/benchmark_workload.cpp) +target_include_directories( + benchmark-workload + PUBLIC common + PRIVATE "${HIREDIS_INCLUDE_DIR}") +target_link_libraries(benchmark-workload PUBLIC CURL::libcurl Threads::Threads + "${HIREDIS_LIBRARY}") + +if(BENCHMARK_SDK STREQUAL "ours") + set(BUILD_TESTING + OFF + CACHE BOOL "" FORCE) + set(BUILD_EXAMPLES + OFF + CACHE BOOL "" FORCE) + set(BUILD_DOCS + OFF + CACHE BOOL "" FORCE) + set(MCP_CPP_SDK_BUILD_CONFORMANCE + OFF + CACHE BOOL "" FORCE) + set(MCP_CPP_SDK_BUILD_SHARED + OFF + CACHE BOOL "" FORCE) + set(MCP_CPP_SDK_BUILD_STATIC + ON + CACHE BOOL "" FORCE) + set(MCP_CPP_SDK_DEFAULT_LINKAGE + static + CACHE STRING "" FORCE) + add_subdirectory("${BENCHMARK_SDK_SOURCE}" sdk) + add_executable(alternative-benchmark-server adapters/ours.cpp) + target_link_libraries(alternative-benchmark-server PRIVATE mcp-cpp-sdk + benchmark-workload) +elseif(BENCHMARK_SDK STREQUAL "hkr04") + set(MCP_BUILD_TESTS + OFF + CACHE BOOL "" FORCE) + set(MCP_MAX_SESSIONS + 0 + CACHE STRING "" FORCE) + set(MCP_SESSION_TIMEOUT + 0 + CACHE STRING "" FORCE) + add_subdirectory("${BENCHMARK_SDK_SOURCE}" sdk) + add_executable(alternative-benchmark-server adapters/hkr04.cpp) + target_include_directories( + alternative-benchmark-server PRIVATE "${BENCHMARK_SDK_SOURCE}/include" + "${BENCHMARK_SDK_SOURCE}/common") + target_link_libraries(alternative-benchmark-server PRIVATE mcp + benchmark-workload) +elseif(BENCHMARK_SDK STREQUAL "fastmcpp") + set(FASTMCPP_BUILD_TESTS + OFF + CACHE BOOL "" FORCE) + set(FASTMCPP_BUILD_EXAMPLES + OFF + CACHE BOOL "" FORCE) + set(FASTMCPP_ENABLE_POST_STREAMING + OFF + CACHE BOOL "" FORCE) + set(FASTMCPP_ENABLE_SAMPLING_HTTP_HANDLERS + OFF + CACHE BOOL "" FORCE) + add_subdirectory("${BENCHMARK_SDK_SOURCE}" sdk) + add_executable(alternative-benchmark-server adapters/fastmcpp.cpp) + target_link_libraries(alternative-benchmark-server PRIVATE fastmcpp_core + benchmark-workload) +elseif(BENCHMARK_SDK STREQUAL "neumann") + set(MCP_BUILD_TESTS + OFF + CACHE BOOL "" FORCE) + set(MCP_BUILD_EXAMPLES + OFF + CACHE BOOL "" FORCE) + set(MCP_WARNINGS_AS_ERRORS + OFF + CACHE BOOL "" FORCE) + set(MCP_ENABLE_HTTP + ON + CACHE BOOL "" FORCE) + set(MCP_HTTP_NO_TLS + ON + CACHE BOOL "" FORCE) + add_subdirectory("${BENCHMARK_SDK_SOURCE}" sdk) + add_executable(alternative-benchmark-server adapters/neumann.cpp) + target_link_libraries(alternative-benchmark-server PRIVATE mcp::http + benchmark-workload) +elseif(BENCHMARK_SDK STREQUAL "cxxmcp") + set(BUILD_TESTING + OFF + CACHE BOOL "" FORCE) + set(CXXMCP_BUILD_SDK + ON + CACHE BOOL "" FORCE) + set(CXXMCP_BUILD_EXAMPLES + OFF + CACHE BOOL "" FORCE) + set(CXXMCP_BUILD_TESTS + OFF + CACHE BOOL "" FORCE) + set(CXXMCP_BUILD_BENCHMARKS + OFF + CACHE BOOL "" FORCE) + set(CXXMCP_BUILD_DOCS + OFF + CACHE BOOL "" FORCE) + set(CXXMCP_ENABLE_HTTP + ON + CACHE BOOL "" FORCE) + set(CXXMCP_ENABLE_OPENSSL + OFF + CACHE BOOL "" FORCE) + set(CXXMCP_ENABLE_AUTH + OFF + CACHE BOOL "" FORCE) + set(CXXMCP_ENABLE_WEBSOCKET + OFF + CACHE BOOL "" FORCE) + add_subdirectory("${BENCHMARK_SDK_SOURCE}" sdk) + add_executable(alternative-benchmark-server adapters/cxxmcp.cpp) + target_link_libraries(alternative-benchmark-server PRIVATE cxxmcp::sdk + benchmark-workload) +else() + message(FATAL_ERROR "No benchmark adapter for '${BENCHMARK_SDK}'") +endif() + +if(CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang") + target_compile_options(alternative-benchmark-server PRIVATE -O3) +endif() diff --git a/benchmark/alternatives/Dockerfile b/benchmark/alternatives/Dockerfile new file mode 100644 index 0000000..2dbba50 --- /dev/null +++ b/benchmark/alternatives/Dockerfile @@ -0,0 +1,30 @@ +FROM ubuntu@sha256:c4a8d5503dfb2a3eb8ab5f807da5bc69a85730fb49b5cfca2330194ebcc41c7b + +ARG DEBIAN_FRONTEND=noninteractive +ARG SDK_NAME +RUN apt-get update && apt-get install -y --no-install-recommends \ + build-essential \ + ca-certificates \ + cmake \ + curl \ + git \ + libcurl4-openssl-dev \ + libhiredis-dev \ + nlohmann-json3-dev \ + pkg-config \ + && rm -rf /var/lib/apt/lists/* + +WORKDIR /src +COPY alternatives/ /src/alternatives/ +COPY alternative-sdks/ /src/alternative-sdks/ + +RUN test -n "${SDK_NAME}" && test -d "/src/alternative-sdks/${SDK_NAME}" +RUN cmake -S /src/alternatives -B /build \ + -DBENCHMARK_SDK="${SDK_NAME}" \ + -DBENCHMARK_SDK_SOURCE="/src/alternative-sdks/${SDK_NAME}" \ + -DCMAKE_BUILD_TYPE=Release \ + && cmake --build /build --target alternative-benchmark-server -j"$(nproc)" + +ENV PORT=8080 +EXPOSE 8080 +CMD ["/build/alternative-benchmark-server"] diff --git a/benchmark/alternatives/Dockerfile.ours b/benchmark/alternatives/Dockerfile.ours new file mode 100644 index 0000000..301cbd9 --- /dev/null +++ b/benchmark/alternatives/Dockerfile.ours @@ -0,0 +1,32 @@ +FROM ubuntu@sha256:c4a8d5503dfb2a3eb8ab5f807da5bc69a85730fb49b5cfca2330194ebcc41c7b + +ARG DEBIAN_FRONTEND=noninteractive +RUN apt-get update && apt-get install -y --no-install-recommends \ + build-essential \ + ca-certificates \ + cmake \ + curl \ + libboost-dev \ + libcurl4-openssl-dev \ + libhiredis-dev \ + libssl-dev \ + nlohmann-json3-dev \ + pkg-config \ + && rm -rf /var/lib/apt/lists/* + +WORKDIR /src +COPY CMakeLists.txt VERSION LICENSE ./ +COPY cmake/ cmake/ +COPY include/ include/ +COPY src/ src/ +COPY benchmark/alternatives/ benchmark/alternatives/ + +RUN cmake -S /src/benchmark/alternatives -B /build \ + -DBENCHMARK_SDK=ours \ + -DBENCHMARK_SDK_SOURCE=/src \ + -DCMAKE_BUILD_TYPE=Release \ + && cmake --build /build --target alternative-benchmark-server -j"$(nproc)" + +ENV PORT=8080 +EXPOSE 8080 +CMD ["/build/alternative-benchmark-server"] diff --git a/benchmark/alternatives/Dockerfile.ours.dockerignore b/benchmark/alternatives/Dockerfile.ours.dockerignore new file mode 100644 index 0000000..372d35b --- /dev/null +++ b/benchmark/alternatives/Dockerfile.ours.dockerignore @@ -0,0 +1,14 @@ +* +!CMakeLists.txt +!VERSION +!LICENSE +!cmake/ +!cmake/** +!include/ +!include/** +!src/ +!src/** +!benchmark/ +benchmark/* +!benchmark/alternatives/ +!benchmark/alternatives/** diff --git a/benchmark/alternatives/adapters/cxxmcp.cpp b/benchmark/alternatives/adapters/cxxmcp.cpp new file mode 100644 index 0000000..d0fbbd8 --- /dev/null +++ b/benchmark/alternatives/adapters/cxxmcp.cpp @@ -0,0 +1,41 @@ +#include "benchmark_workload.hpp" + +#include +#include + +#include +#include +#include + +namespace { + +int port() { + const char* value = std::getenv("PORT"); + return value ? std::stoi(value) : 8080; +} + +} // namespace + +int main() { + using Json = mcp::protocol::Json; + + mcp_benchmark::Workload workload("cpp-sdk"); + auto server = mcp::ServerPeer::builder(); + server.name("benchmark-cpp-sdk").version("1.0.0").streamable_http("0.0.0.0", port(), "/mcp"); + + const auto add_tool = [&](const std::string& name, const std::string& description) { + mcp::protocol::ToolDefinition definition; + definition.name = name; + definition.description = description; + definition.input_schema = Json::parse(mcp_benchmark::Workload::input_schema(name)); + server.tool( + std::move(definition), [&workload, name](const Json& args) { + return mcp::protocol::ToolResult::text(workload.invoke(name, args.dump())); + }); + }; + add_tool("search_products", "Search products and merge popularity data"); + add_tool("get_user_cart", "Get a cart with recent order history"); + add_tool("checkout", "Calculate and record a checkout"); + + return server.run(); +} diff --git a/benchmark/alternatives/adapters/fastmcpp.cpp b/benchmark/alternatives/adapters/fastmcpp.cpp new file mode 100644 index 0000000..80839eb --- /dev/null +++ b/benchmark/alternatives/adapters/fastmcpp.cpp @@ -0,0 +1,48 @@ +#include "benchmark_workload.hpp" + +#include "fastmcpp/app.hpp" +#include "fastmcpp/mcp/handler.hpp" +#include "fastmcpp/server/streamable_http_server.hpp" + +#include +#include +#include +#include +#include + +namespace { + +int port() { + const char* value = std::getenv("PORT"); + return value ? std::stoi(value) : 8080; +} + +} // namespace + +int main() { + using fastmcpp::Json; + mcp_benchmark::Workload workload("cpp-sdk"); + fastmcpp::FastMCP app("benchmark-cpp-sdk", "1.0.0"); + + const auto add_tool = [&](const std::string& name, const std::string& description) { + fastmcpp::FastMCP::ToolOptions options; + options.description = description; + app.tool( + name, Json::parse(mcp_benchmark::Workload::input_schema(name)), + [&workload, name](const Json& args) { return workload.invoke(name, args.dump()); }, + std::move(options)); + }; + add_tool("search_products", "Search products and merge popularity data"); + add_tool("get_user_cart", "Get a cart with recent order history"); + add_tool("checkout", "Calculate and record a checkout"); + + auto handler = fastmcpp::mcp::make_mcp_handler(app); + fastmcpp::server::StreamableHttpServerWrapper server(std::move(handler), "0.0.0.0", port(), "/mcp"); + if (!server.start()) { + std::cerr << "failed to start FastMCPP benchmark server\n"; + return EXIT_FAILURE; + } + while (true) { + std::this_thread::sleep_for(std::chrono::hours(24)); + } +} diff --git a/benchmark/alternatives/adapters/hkr04.cpp b/benchmark/alternatives/adapters/hkr04.cpp new file mode 100644 index 0000000..8fbd649 --- /dev/null +++ b/benchmark/alternatives/adapters/hkr04.cpp @@ -0,0 +1,49 @@ +#include "benchmark_workload.hpp" + +#include "mcp_server.h" +#include "mcp_tool.h" + +#include +#include + +namespace { + +int port() { + const char* value = std::getenv("PORT"); + return value ? std::stoi(value) : 8080; +} + +mcp::tool make_tool(const std::string& name, const std::string& description) { + return mcp::tool{name, description, mcp::json::parse(mcp_benchmark::Workload::input_schema(name)), + mcp::json::object()}; +} + +} // namespace + +int main() { + mcp::set_log_level(mcp::log_level::error); + mcp_benchmark::Workload workload("cpp-sdk"); + mcp::server::configuration config; + config.host = "0.0.0.0"; + config.port = port(); + config.mcp_endpoint = "/mcp"; + config.max_sessions = 0; + config.session_timeout = 0; + // This separate async pool serves legacy paths, not the synchronous + // Streamable HTTP endpoint measured by this benchmark. + config.threadpool_size = 2; + mcp::server server(config); + server.set_server_info("benchmark-cpp-sdk", "1.0.0"); + server.set_capabilities({{"tools", mcp::json::object()}}); + + const auto register_tool = [&](const std::string& name, const std::string& description) { + server.register_tool(make_tool(name, description), [&workload, name](const mcp::json& args, + const std::string&) { + return mcp::json::array({{{"type", "text"}, {"text", workload.invoke(name, args.dump())}}}); + }); + }; + register_tool("search_products", "Search products and merge popularity data"); + register_tool("get_user_cart", "Get a cart with recent order history"); + register_tool("checkout", "Calculate and record a checkout"); + return server.start(true) ? EXIT_SUCCESS : EXIT_FAILURE; +} diff --git a/benchmark/alternatives/adapters/neumann.cpp b/benchmark/alternatives/adapters/neumann.cpp new file mode 100644 index 0000000..f7c2e68 --- /dev/null +++ b/benchmark/alternatives/adapters/neumann.cpp @@ -0,0 +1,54 @@ +#include "benchmark_workload.hpp" + +#include "mcp/http_server_host.hpp" +#include "mcp/mcp.hpp" + +#include +#include +#include +#include + +namespace { + +int port() { + const char* value = std::getenv("PORT"); + return value ? std::stoi(value) : 8080; +} + +mcp::CallToolResult text_result(std::string text) { + return mcp::CallToolResult{ + .content = {mcp::TextContent{.text = std::move(text)}}, + .is_error = false, + }; +} + +} // namespace + +int main() { + mcp::set_log_level(mcp::LogLevel::off); + mcp_benchmark::Workload workload("cpp-sdk"); + mcp::HttpServerHost::Options options; + options.host = "0.0.0.0"; + options.port = port(); + options.path = "/mcp"; + + mcp::HttpServerHost host( + mcp::Implementation{.name = "benchmark-cpp-sdk", .version = "1.0.0"}, std::move(options), + [&workload](mcp::Server& server) { + const auto add_tool = [&](const std::string& name, const std::string& description) { + server.tool( + name, nlohmann::json::parse(mcp_benchmark::Workload::input_schema(name)), + [&workload, name](const nlohmann::json& args) { + return text_result(workload.invoke(name, args.dump())); + }, + std::nullopt, description); + }; + add_tool("search_products", "Search products and merge popularity data"); + add_tool("get_user_cart", "Get a cart with recent order history"); + add_tool("checkout", "Calculate and record a checkout"); + }); + host.start(); + while (true) { + std::this_thread::sleep_for(std::chrono::hours(24)); + } +} diff --git a/benchmark/alternatives/adapters/ours.cpp b/benchmark/alternatives/adapters/ours.cpp new file mode 100644 index 0000000..9175a64 --- /dev/null +++ b/benchmark/alternatives/adapters/ours.cpp @@ -0,0 +1,69 @@ +#include "benchmark_workload.hpp" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { + +constexpr std::size_t kHandlerThreadCount = MCP_BENCHMARK_REQUEST_CONCURRENCY; +constexpr std::size_t kIoThreadCount = 2; + +unsigned short port() { + const char* value = std::getenv("PORT"); + return value ? static_cast(std::stoi(value)) : 8080; +} + +std::unique_ptr make_server(const std::shared_ptr& workload) { + mcp::ServerCapabilities capabilities; + capabilities.tools = mcp::ServerCapabilities::ToolsCapability{}; + auto server = std::make_unique(mcp::Implementation{"benchmark-cpp-sdk", "1.0.0"}, + std::move(capabilities)); + + const auto add_tool = [&](const std::string& name, const std::string& description) { + server->add_tool( + name, description, nlohmann::json::parse(mcp_benchmark::Workload::input_schema(name)), + [workload, name](const nlohmann::json& args) { + return mcp::make_tool_text_result(workload->invoke(name, args.dump())); + }); + }; + add_tool("search_products", "Search products and merge popularity data"); + add_tool("get_user_cart", "Get a cart with recent order history"); + add_tool("checkout", "Calculate and record a checkout"); + return server; +} + +} // namespace + +int main() { + namespace asio = boost::asio; + + auto workload = std::make_shared("cpp-sdk"); + asio::io_context io_context; + asio::thread_pool handler_pool(kHandlerThreadCount); + + mcp::StreamableHttpSessionManager manager( + io_context.get_executor(), "0.0.0.0", port(), + [workload](const asio::any_io_executor&) { return make_server(workload); }); + manager.set_tool_executor(handler_pool.get_executor()); + + const auto io_work = asio::make_work_guard(io_context); + asio::co_spawn(io_context, manager.listen(), asio::detached); + + std::vector io_threads; + io_threads.reserve(kIoThreadCount); + for (std::size_t i = 0; i < kIoThreadCount; ++i) { + io_threads.emplace_back([&io_context] { io_context.run(); }); + } + for (auto& thread : io_threads) { + thread.join(); + } +} diff --git a/benchmark/alternatives/benchmark.js b/benchmark/alternatives/benchmark.js new file mode 100644 index 0000000..55ab98d --- /dev/null +++ b/benchmark/alternatives/benchmark.js @@ -0,0 +1,567 @@ +import {check, sleep} from 'k6'; +import http from 'k6/http'; +import {Counter, Rate, Trend} from 'k6/metrics'; + +const SERVER_URL = __ENV.SERVER_URL || 'http://localhost:8080/mcp'; +const SERVER_NAME = __ENV.SERVER_NAME || 'unknown'; +const PROTOCOL_VERSION = __ENV.MCP_PROTOCOL_VERSION || '2024-11-05'; +const EXPECTED_PROTOCOL_VERSION = __ENV.EXPECTED_PROTOCOL_VERSION || null; +const MODE = (__ENV.BENCHMARK_MODE || 'measurement').toLowerCase(); +const BENCHMARK_CONTRACT = __ENV.BENCHMARK_CONTRACT || 'upstream-v2-strict-mcp-v1'; + +const VUS = Number.parseInt(__ENV.BENCHMARK_VUS || '50', 10); +const WARMUP_RAMP_DURATION = __ENV.BENCHMARK_RAMP_DURATION || '15s'; +const WARMUP_LOAD_DURATION = __ENV.BENCHMARK_WARMUP_DURATION || '60s'; +const MEASUREMENT_DURATION = __ENV.BENCHMARK_MEASURE_DURATION || '5m'; + +if (!Number.isInteger(VUS) || VUS <= 0) { + throw new Error(`BENCHMARK_VUS must be a positive integer, got "${__ENV.BENCHMARK_VUS}"`); +} + +if (MODE !== 'warmup' && MODE !== 'measurement') { + throw new Error( + `BENCHMARK_MODE must be "warmup" or "measurement", got "${MODE}"`, + ); +} + +if (BENCHMARK_CONTRACT !== 'upstream-v2-strict-mcp-v1') { + throw new Error( + 'BENCHMARK_CONTRACT must be "upstream-v2-strict-mcp-v1", ' + + `got "${BENCHMARK_CONTRACT}"`, + ); +} + +const scenario = MODE === 'warmup' ? + { + executor: 'ramping-vus', + startVUs: 0, + stages: [ + {duration: WARMUP_RAMP_DURATION, target: VUS}, + {duration: WARMUP_LOAD_DURATION, target: VUS}, + ], + gracefulRampDown: '0s', + gracefulStop: '0s', + } : + { + executor: 'constant-vus', + vus: VUS, + duration: MEASUREMENT_DURATION, + // Do not add a server-dependent shutdown tail to the Counter rate + // denominator. Measurement is exactly the configured constant-VU window. + gracefulStop: '0s', + }; + +export const options = { + scenarios: { + workload: scenario, + }, + thresholds: { + checks: [{threshold: 'rate==1', abortOnFail: true}], + http_req_failed: [{threshold: 'rate==0', abortOnFail: true}], + mcp_error_rate: [{threshold: 'rate==0', abortOnFail: true}], + mcp_errors: [{threshold: 'count==0', abortOnFail: true}], + }, + summaryTrendStats: [ + 'avg', + 'min', + 'med', + 'max', + 'p(90)', + 'p(95)', + 'p(99)', + 'count', + ], +}; + +const initializeDuration = new Trend('mcp_initialize_duration', true); +const toolsListDuration = new Trend('mcp_tools_list_duration', true); +const searchProductsDuration = new Trend('mcp_search_products_duration', true); +const getUserCartDuration = new Trend('mcp_get_user_cart_duration', true); +const checkoutDuration = new Trend('mcp_checkout_duration', true); +const combinedToolDuration = new Trend('mcp_tool_duration', true); +const sessionDuration = new Trend('mcp_session_duration', true); + +// A benchmark operation is one tools/call or tools/list request whose shape +// and selected-contract checks have run. Lifecycle traffic is reported separately so +// headline throughput does not conflate useful work with session management. +const benchmarkOperations = new Counter('benchmark_operations'); +const mcpMessages = new Counter('mcp_messages'); +const sessionsStarted = new Counter('mcp_sessions_started'); +const sessionsCompleted = new Counter('mcp_sessions_completed'); +const sessionsSuccessful = new Counter('mcp_sessions_successful'); +const sessionsFailed = new Counter('mcp_sessions_failed'); +const mcpErrors = new Counter('mcp_errors'); +const mcpErrorRate = new Rate('mcp_error_rate'); + +const BASE_HEADERS = { + Accept: 'application/json, text/event-stream', + 'Content-Type': 'application/json', +}; + +const EXPECTED_TOOLS = ['search_products', 'get_user_cart', 'checkout']; +const SUPPORTED_PROTOCOL_VERSIONS = new Set([ + '2024-11-05', + '2025-03-26', + '2025-06-18', + '2025-11-25', +]); + +function responseHeader(response, name) { + const wanted = name.toLowerCase(); + for (const [key, value] of Object.entries(response.headers || {})) { + if (key.toLowerCase() === wanted) + return value; + } + return null; +} + +function parseJson(text) { + try { + return JSON.parse(text); + } catch (_) { + return null; + } +} + +function parseRpcMessage(body, expectedId, mediaType) { + if (!body || !body.trim()) + return null; + + let candidates = []; + if (mediaType === 'application/json') { + const direct = parseJson(body.trim()); + // This harness sends one non-batch JSON-RPC request at a time. A batch + // or any extra response would make the measured exchange non-equivalent. + if (direct === null || Array.isArray(direct) || typeof direct !== 'object') { + return null; + } + candidates = [direct]; + } else if (mediaType === 'text/event-stream') { + const events = body.replace(/\r\n/g, '\n').split('\n\n'); + for (const event of events) { + const data = event.split('\n') + .filter((line) => line.startsWith('data:')) + .map((line) => line.slice(5).trimStart()) + .join('\n'); + if (!data) + continue; + const message = parseJson(data); + if (message === null || Array.isArray(message) || typeof message !== 'object') { + return null; + } + candidates.push(message); + } + } else { + return null; + } + + if (candidates.length !== 1 || candidates[0].id !== expectedId) + return null; + return candidates[0]; +} + +function mcpResponseMediaType(response) { + const contentType = responseHeader(response, 'Content-Type'); + if (typeof contentType !== 'string') + return null; + const mediaType = contentType.split(';', 1)[0].trim().toLowerCase(); + return mediaType === 'application/json' || mediaType === 'text/event-stream' ? mediaType : null; +} + +function recordCheck(label, passed) { + const checks = {}; + checks[label] = (value) => value === true; + return check(passed, checks); +} + +function initializeSession() { + const id = 1; + mcpMessages.add(1); + const response = http.post( + SERVER_URL, + JSON.stringify({ + jsonrpc: '2.0', + id, + method: 'initialize', + params: { + protocolVersion: PROTOCOL_VERSION, + capabilities: {}, + clientInfo: {name: 'mcp-cpp-sdk-benchmark', version: '1.0'}, + }, + }), + { + headers: BASE_HEADERS, + timeout: '30s', + tags: {name: 'mcp_initialize'}, + }, + ); + initializeDuration.add(response.timings.duration); + + const mediaType = mcpResponseMediaType(response); + const message = parseRpcMessage(response.body, id, mediaType); + const result = message && message.result; + const valid = response.status === 200 && mediaType !== null && message !== null && + message.jsonrpc === '2.0' && !message.error && result && + SUPPORTED_PROTOCOL_VERSIONS.has(result.protocolVersion) && + (EXPECTED_PROTOCOL_VERSION === null || result.protocolVersion === EXPECTED_PROTOCOL_VERSION) && + result.capabilities && typeof result.capabilities === 'object' && result.capabilities.tools && + typeof result.capabilities.tools === 'object' && result.serverInfo && + typeof result.serverInfo === 'object' && typeof result.serverInfo.name === 'string' && + result.serverInfo.name.length > 0 && typeof result.serverInfo.version === 'string' && + result.serverInfo.version.length > 0; + recordCheck('initialize returns a valid negotiated result', valid); + + return { + response, + valid, + sessionId: responseHeader(response, 'Mcp-Session-Id'), + protocolVersion: valid ? result.protocolVersion : PROTOCOL_VERSION, + }; +} + +function postInitializeHeaders(context) { + const headers = Object.assign({}, BASE_HEADERS, { + 'MCP-Protocol-Version': context.protocolVersion, + }); + if (context.sessionId) + headers['Mcp-Session-Id'] = context.sessionId; + return headers; +} + +function sendInitialized(context) { + mcpMessages.add(1); + const response = http.post( + SERVER_URL, + JSON.stringify({jsonrpc: '2.0', method: 'notifications/initialized'}), + { + headers: postInitializeHeaders(context), + timeout: '5s', + tags: {name: 'mcp_initialized_notification'}, + }, + ); + const valid = response.status === 202 && (!response.body || response.body.trim() === ''); + recordCheck('initialized notification is accepted', valid); + return valid; +} + +function requestOperation(context, method, params, metricName) { + const id = 2; + mcpMessages.add(1); + const response = http.post( + SERVER_URL, + JSON.stringify({jsonrpc: '2.0', id, method, params}), + { + headers: postInitializeHeaders(context), + timeout: '30s', + tags: {name: metricName}, + }, + ); + return { + message: parseRpcMessage( + response.body, + id, + mcpResponseMediaType(response), + ), + response, + }; +} + +function closeSession(context) { + if (!context.sessionId) + return true; + + const response = http.del(SERVER_URL, null, { + headers: postInitializeHeaders(context), + timeout: '5s', + tags: {name: 'mcp_delete_session'}, + // MCP permits 405 when a server does not offer client-initiated session + // termination. Treat that specified response as transport-successful. + responseCallback: http.expectedStatuses({min: 200, max: 299}, 405), + }); + const valid = (response.status >= 200 && response.status < 300) || response.status === 405; + recordCheck('session DELETE is accepted', valid); + return valid; +} + +function finishSession(startedAt, failed) { + sessionsCompleted.add(1); + sessionDuration.add(Date.now() - startedAt); + mcpErrorRate.add(failed); + mcpErrors.add(failed ? 1 : 0); + if (failed) { + sessionsFailed.add(1); + } else { + sessionsSuccessful.add(1); + } +} + +function validRpcResult(response, message) { + return response.status === 200 && mcpResponseMediaType(response) !== null && message !== null && + message.jsonrpc === '2.0' && !message.error && message.result && + typeof message.result === 'object'; +} + +function callTool(tool) { + sessionsStarted.add(1); + const startedAt = Date.now(); + const context = initializeSession(); + let failed = !context.valid; + + if (!sendInitialized(context)) + failed = true; + + const operation = requestOperation( + context, + 'tools/call', + {name: tool.name, arguments: tool.arguments}, + `mcp_tool_${tool.name}`, + ); + const shapeValid = validRpcResult(operation.response, operation.message) && + operation.message.result.isError !== true && Array.isArray(operation.message.result.content); + const contractValid = shapeValid && tool.validateContract(operation.message.result); + recordCheck(`${tool.name} returns a valid tool result`, shapeValid); + recordCheck(`${tool.name} satisfies the benchmark contract`, contractValid); + if (!shapeValid || !contractValid) + failed = true; + tool.duration.add(operation.response.timings.duration); + combinedToolDuration.add(operation.response.timings.duration); + benchmarkOperations.add(1); + + if (!closeSession(context)) + failed = true; + finishSession(startedAt, failed); +} + +function listTools() { + sessionsStarted.add(1); + const startedAt = Date.now(); + const context = initializeSession(); + let failed = !context.valid; + + if (!sendInitialized(context)) + failed = true; + + const operation = requestOperation( + context, + 'tools/list', + {}, + 'mcp_tools_list', + ); + const resultValid = validRpcResult(operation.response, operation.message) && + Array.isArray(operation.message.result.tools); + const tools = resultValid ? operation.message.result.tools : []; + const names = tools.map((tool) => tool.name); + const expectedNamesPresent = resultValid && tools.length === EXPECTED_TOOLS.length && + new Set(names).size === EXPECTED_TOOLS.length && + EXPECTED_TOOLS.every((name) => names.includes(name)); + const contractValid = expectedNamesPresent && + tools.every( + (tool) => tool.inputSchema && typeof tool.inputSchema === 'object' && + !Array.isArray(tool.inputSchema), + ); + recordCheck('tools/list returns a valid tool collection', resultValid); + recordCheck('tools/list satisfies the benchmark contract', contractValid); + if (!resultValid || !contractValid) + failed = true; + toolsListDuration.add(operation.response.timings.duration); + benchmarkOperations.add(1); + + if (!closeSession(context)) + failed = true; + finishSession(startedAt, failed); +} + +function textContent(result) { + if (!result || !Array.isArray(result.content)) + return null; + const block = result.content.find( + (content) => content && content.type === 'text' && typeof content.text === 'string', + ); + return block ? parseJson(block.text) : null; +} + +function toolsForUser(userId) { + return [ + { + name: 'search_products', + arguments: { + category: 'Electronics', + min_price: 50.0, + max_price: 500.0, + limit: 10, + }, + duration: searchProductsDuration, + validateContract: (result) => { + const value = textContent(result); + return value !== null && value.total_found === 2251 && Array.isArray(value.products) && + value.products.length === 10 && Array.isArray(value.top10_popular_ids) && + value.top10_popular_ids.length === 10; + }, + }, + { + name: 'get_user_cart', + arguments: {user_id: userId}, + duration: getUserCartDuration, + validateContract: (result) => { + const value = textContent(result); + return value !== null && value.user_id === userId && value.cart && + Array.isArray(value.cart.items) && value.cart.items.length >= 1 && + Array.isArray(value.recent_history) && value.recent_history.length === 5; + }, + }, + { + name: 'checkout', + arguments: { + user_id: userId, + items: [ + {product_id: 42, quantity: 2}, + {product_id: 1337, quantity: 1}, + ], + }, + duration: checkoutDuration, + validateContract: (result) => { + const value = textContent(result); + return value !== null && value.user_id === userId && value.status === 'confirmed' && + typeof value.total === 'number' && value.total > 0 && value.items_count === 2 && + typeof value.rate_limit_count === 'number'; + }, + }, + ]; +} + +export default function() { + const userNumber = ((__VU - 1) % 1000) + 1; + const userId = `user-${String(userNumber).padStart(5, '0')}`; + + for (const tool of toolsForUser(userId)) + callTool(tool); + listTools(); + sleep(0.05); +} + +function metricValue(data, metricName, valueName, fallback = 0) { + const metric = data.metrics[metricName]; + if (!metric || metric.values[valueName] === undefined) + return fallback; + return metric.values[valueName]; +} + +function trend(data, metricName) { + const metric = data.metrics[metricName]; + if (!metric) + return null; + return { + count: metric.values.count, + avg_ms: metric.values.avg, + min_ms: metric.values.min, + p50_ms: metric.values.med, + p90_ms: metric.values['p(90)'], + p95_ms: metric.values['p(95)'], + p99_ms: metric.values['p(99)'], + max_ms: metric.values.max, + }; +} + +function counter(data, metricName) { + return { + count: metricValue(data, metricName, 'count'), + per_second: metricValue(data, metricName, 'rate'), + }; +} + +function actualDurationMs(data) { + if (data.state && typeof data.state.testRunDurationMs === 'number') { + return data.state.testRunDurationMs; + } + return null; +} + +function checkBreakdown(group, output = {}) { + if (!group || typeof group !== 'object') + return output; + for (const item of group.checks || []) { + output[item.name] = {passes: item.passes || 0, fails: item.fails || 0}; + } + for (const child of group.groups || []) + checkBreakdown(child, output); + return output; +} + +export function handleSummary(data) { + const outputPath = __ENV.OUTPUT_PATH || `${SERVER_NAME}_${MODE}_summary.json`; + const durationMs = actualDurationMs(data); + const summary = { + schema_version: 2, + server: SERVER_NAME, + timestamp: new Date().toISOString(), + config: { + mode: MODE, + server_url: SERVER_URL, + requested_protocol_version: PROTOCOL_VERSION, + expected_negotiated_protocol_version: EXPECTED_PROTOCOL_VERSION, + eligibility_contract: BENCHMARK_CONTRACT, + vus: VUS, + executor: scenario.executor, + ramp_duration: MODE === 'warmup' ? WARMUP_RAMP_DURATION : null, + warmup_load_duration: MODE === 'warmup' ? WARMUP_LOAD_DURATION : null, + measurement_duration: MODE === 'measurement' ? MEASUREMENT_DURATION : null, + configured_duration: MODE === 'measurement' ? + MEASUREMENT_DURATION : + `${WARMUP_RAMP_DURATION} ramp + ${WARMUP_LOAD_DURATION} load`, + actual_duration_seconds: durationMs === null ? null : durationMs / 1000, + }, + rates: { + operations: counter(data, 'benchmark_operations'), + mcp_messages: counter(data, 'mcp_messages'), + raw_http: counter(data, 'http_reqs'), + }, + latency: { + combined_tool_call: trend(data, 'mcp_tool_duration'), + full_session: trend(data, 'mcp_session_duration'), + initialize: trend(data, 'mcp_initialize_duration'), + tools_list: trend(data, 'mcp_tools_list_duration'), + tools: { + search_products: trend(data, 'mcp_search_products_duration'), + get_user_cart: trend(data, 'mcp_get_user_cart_duration'), + checkout: trend(data, 'mcp_checkout_duration'), + }, + }, + sessions: { + started: metricValue(data, 'mcp_sessions_started', 'count'), + completed: metricValue(data, 'mcp_sessions_completed', 'count'), + successful: metricValue(data, 'mcp_sessions_successful', 'count'), + failed: metricValue(data, 'mcp_sessions_failed', 'count'), + }, + check_breakdown: checkBreakdown(data.root_group), + errors: { + mcp: metricValue(data, 'mcp_errors', 'count'), + mcp_rate: metricValue(data, 'mcp_error_rate', 'rate'), + http: metricValue(data, 'http_req_failed', 'passes'), + http_rate: metricValue(data, 'http_req_failed', 'rate'), + checks: metricValue(data, 'checks', 'fails'), + check_pass_rate: metricValue(data, 'checks', 'rate'), + }, + }; + + const operations = summary.rates.operations; + const rawHttp = summary.rates.raw_http; + const combined = summary.latency.combined_tool_call; + const latencyLine = combined === null ? 'tool latency: no samples' : + `tool latency: p50=${combined.p50_ms.toFixed(2)}ms ` + + `p95=${combined.p95_ms.toFixed(2)}ms p99=${combined.p99_ms.toFixed(2)}ms`; + const stdout = [ + '', + `${SERVER_NAME} ${MODE} benchmark`, + `actual duration: ${summary.config.actual_duration_seconds ?? 'unknown'}s`, + `operations: ${operations.count} (${operations.per_second.toFixed(2)}/s)`, + `raw HTTP requests: ${rawHttp.count} (${rawHttp.per_second.toFixed(2)}/s)`, + latencyLine, + `errors: MCP=${summary.errors.mcp} HTTP=${summary.errors.http} checks=${summary.errors.checks}`, + '', + ].join('\n'); + + return { + [outputPath]: JSON.stringify(summary, null, 2), + stdout, + }; +} diff --git a/benchmark/alternatives/common/benchmark_workload.cpp b/benchmark/alternatives/common/benchmark_workload.cpp new file mode 100644 index 0000000..178b88d --- /dev/null +++ b/benchmark/alternatives/common/benchmark_workload.cpp @@ -0,0 +1,452 @@ +#include "benchmark_workload.hpp" + +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mcp_benchmark { +namespace { + +using Json = nlohmann::json; + +std::string env_or(const char* name, const char* fallback) { + const char* value = std::getenv(name); + return value && *value ? value : fallback; +} + +class ThreadPool { + public: + explicit ThreadPool(std::size_t count) { + workers_.reserve(count); + for (std::size_t i = 0; i < count; ++i) { + workers_.emplace_back([this] { + for (;;) { + std::function task; + { + std::unique_lock lock(mutex_); + cv_.wait(lock, [this] { return stopping_ || !tasks_.empty(); }); + if (stopping_ && tasks_.empty()) { + return; + } + task = std::move(tasks_.front()); + tasks_.pop(); + } + task(); + } + }); + } + } + + ~ThreadPool() { + { + std::lock_guard lock(mutex_); + stopping_ = true; + } + cv_.notify_all(); + for (auto& worker : workers_) { + worker.join(); + } + } + + template + auto submit(F&& fn) -> std::future { + using Result = decltype(fn()); + auto task = std::make_shared>(std::forward(fn)); + auto result = task->get_future(); + { + std::lock_guard lock(mutex_); + tasks_.emplace([task] { (*task)(); }); + } + cv_.notify_one(); + return result; + } + + private: + std::mutex mutex_; + std::condition_variable cv_; + std::queue> tasks_; + std::vector workers_; + bool stopping_ = false; +}; + +std::size_t write_body(char* data, std::size_t size, std::size_t count, void* opaque) { + auto* output = static_cast(opaque); + output->append(data, size * count); + return size * count; +} + +class HttpClient { + public: + explicit HttpClient(std::string base_url) : base_url_(std::move(base_url)) { + while (!base_url_.empty() && base_url_.back() == '/') { + base_url_.pop_back(); + } + } + + std::string get(const std::string& path) const { return request(base_url_ + path, nullptr); } + + std::string post(const std::string& path, const std::string& body) const { + return request(base_url_ + path, &body); + } + + private: + static std::string request(const std::string& url, const std::string* body) { + thread_local std::unique_ptr curl(curl_easy_init(), + &curl_easy_cleanup); + if (!curl) { + throw std::runtime_error("curl_easy_init failed"); + } + + std::string output; + curl_easy_reset(curl.get()); + curl_easy_setopt(curl.get(), CURLOPT_URL, url.c_str()); + curl_easy_setopt(curl.get(), CURLOPT_WRITEFUNCTION, write_body); + curl_easy_setopt(curl.get(), CURLOPT_WRITEDATA, &output); + curl_easy_setopt(curl.get(), CURLOPT_TIMEOUT_MS, 10000L); + curl_easy_setopt(curl.get(), CURLOPT_CONNECTTIMEOUT_MS, 2000L); + curl_easy_setopt(curl.get(), CURLOPT_TCP_KEEPALIVE, 1L); + curl_easy_setopt(curl.get(), CURLOPT_NOSIGNAL, 1L); + + curl_slist* headers = nullptr; + if (body) { + headers = curl_slist_append(headers, "Content-Type: application/json"); + curl_easy_setopt(curl.get(), CURLOPT_HTTPHEADER, headers); + curl_easy_setopt(curl.get(), CURLOPT_POST, 1L); + curl_easy_setopt(curl.get(), CURLOPT_POSTFIELDS, body->data()); + curl_easy_setopt(curl.get(), CURLOPT_POSTFIELDSIZE, static_cast(body->size())); + } + + const CURLcode status = curl_easy_perform(curl.get()); + long response_code = 0; + curl_easy_getinfo(curl.get(), CURLINFO_RESPONSE_CODE, &response_code); + if (headers) { + curl_slist_free_all(headers); + } + if (status != CURLE_OK) { + throw std::runtime_error(std::string("HTTP request failed: ") + curl_easy_strerror(status)); + } + if (response_code < 200 || response_code >= 300) { + throw std::runtime_error("HTTP request returned status " + std::to_string(response_code)); + } + return output; + } + + std::string base_url_; +}; + +struct RedisEndpoint { + std::string host = "redis"; + int port = 6379; +}; + +RedisEndpoint parse_redis_url(std::string url) { + constexpr const char* prefix = "redis://"; + if (url.rfind(prefix, 0) == 0) { + url.erase(0, std::char_traits::length(prefix)); + } + const auto slash = url.find('/'); + if (slash != std::string::npos) { + url.resize(slash); + } + const auto colon = url.rfind(':'); + if (colon == std::string::npos) { + return {url, 6379}; + } + return {url.substr(0, colon), std::stoi(url.substr(colon + 1))}; +} + +class RedisClient { + public: + explicit RedisClient(RedisEndpoint endpoint) : endpoint_(std::move(endpoint)) {} + + std::vector string_array(const std::vector& args) const { + Reply reply = command(args); + std::vector values; + if (!reply.value || reply.value->type != REDIS_REPLY_ARRAY) { + return values; + } + values.reserve(reply.value->elements); + for (std::size_t i = 0; i < reply.value->elements; ++i) { + const redisReply* item = reply.value->element[i]; + values.emplace_back(item && item->str ? std::string(item->str, item->len) : ""); + } + return values; + } + + std::int64_t integer(const std::vector& args) const { + Reply reply = command(args); + return reply.value && reply.value->type == REDIS_REPLY_INTEGER ? reply.value->integer : 0; + } + + void discard(const std::vector& args) const { (void)command(args); } + + private: + struct Reply { + redisReply* value = nullptr; + ~Reply() { + if (value) { + freeReplyObject(value); + } + } + Reply(const Reply&) = delete; + Reply& operator=(const Reply&) = delete; + Reply(Reply&& other) noexcept : value(std::exchange(other.value, nullptr)) {} + explicit Reply(redisReply* input) : value(input) {} + }; + + struct Connection { + redisContext* value = nullptr; + ~Connection() { + if (value) { + redisFree(value); + } + } + }; + + Reply command(const std::vector& args) const { + thread_local Connection connection; + if (!connection.value || connection.value->err) { + if (connection.value) { + redisFree(connection.value); + } + connection.value = redisConnect(endpoint_.host.c_str(), endpoint_.port); + if (!connection.value || connection.value->err) { + throw std::runtime_error("Redis connection failed"); + } + } + + std::vector argv; + std::vector lengths; + argv.reserve(args.size()); + lengths.reserve(args.size()); + for (const auto& arg : args) { + argv.push_back(arg.data()); + lengths.push_back(arg.size()); + } + auto* raw = static_cast(redisCommandArgv( + connection.value, static_cast(argv.size()), argv.data(), lengths.data())); + if (!raw) { + redisFree(connection.value); + connection.value = nullptr; + throw std::runtime_error("Redis command failed"); + } + return Reply(raw); + } + + RedisEndpoint endpoint_; +}; + +std::string number(double value) { + std::ostringstream out; + out << std::setprecision(15) << value; + return out.str(); +} + +int user_number(const std::string& user_id) { + const auto dash = user_id.rfind('-'); + if (dash == std::string::npos) { + return 42; + } + try { + return std::stoi(user_id.substr(dash + 1)); + } catch (...) { + return 42; + } +} + +} // namespace + +class Workload::Impl { + public: + explicit Impl(std::string server_type) + : server_type_(std::move(server_type)), + http_(env_or("API_SERVICE_URL", "http://api-service:8100")), + redis_(parse_redis_url(env_or("REDIS_URL", "redis://redis:6379"))), + pool_(64) { + static const int curl_initialized = [] { return curl_global_init(CURL_GLOBAL_DEFAULT); }(); + if (curl_initialized != CURLE_OK) { + throw std::runtime_error("curl_global_init failed"); + } + } + + std::string invoke(const std::string& name, const std::string& raw_arguments) { + const Json args = raw_arguments.empty() ? Json::object() : Json::parse(raw_arguments); + if (name == "search_products") { + return search_products(args).dump(); + } + if (name == "get_user_cart") { + return get_user_cart(args).dump(); + } + if (name == "checkout") { + return checkout(args).dump(); + } + throw std::invalid_argument("unknown benchmark tool: " + name); + } + + private: + Json search_products(const Json& args) { + const std::string category = args.value("category", "Electronics"); + const double min_price = args.value("min_price", 50.0); + const double max_price = args.value("max_price", 500.0); + const int limit = args.value("limit", 10); + const std::string path = "/products/search?category=" + category + + "&min_price=" + number(min_price) + "&max_price=" + number(max_price) + + "&limit=" + std::to_string(limit); + + auto search = pool_.submit([this, path] { return http_.get(path); }); + auto popular = pool_.submit( + [this] { return redis_.string_array({"ZREVRANGE", "bench:popular", "0", "9"}); }); + Json search_data = Json::parse(search.get()); + const auto popular_raw = popular.get(); + + Json ids = Json::array(); + std::vector top_ids; + for (const auto& member : popular_raw) { + const auto colon = member.find(':'); + if (colon == std::string::npos) { + continue; + } + const int id = std::stoi(member.substr(colon + 1)); + top_ids.push_back(id); + ids.push_back(id); + } + + Json products = Json::array(); + for (const auto& product : search_data.value("products", Json::array())) { + const int id = product.value("id", 0); + const auto it = std::find(top_ids.begin(), top_ids.end(), id); + const int rank = it == top_ids.end() ? 0 : static_cast(it - top_ids.begin()) + 1; + products.push_back({{"id", id}, + {"sku", product.value("sku", "")}, + {"name", product.value("name", "")}, + {"price", product.value("price", 0.0)}, + {"rating", product.value("rating", 0.0)}, + {"popularity_rank", rank}}); + } + return {{"category", category}, + {"total_found", search_data.value("total_found", 0)}, + {"products", std::move(products)}, + {"top10_popular_ids", std::move(ids)}, + {"server_type", server_type_}}; + } + + Json get_user_cart(const Json& args) { + const std::string user_id = args.value("user_id", "user-00042"); + const auto hash = redis_.string_array({"HGETALL", "bench:cart:" + user_id}); + Json items = Json::array(); + double total = 0.0; + for (std::size_t i = 0; i + 1 < hash.size(); i += 2) { + if (hash[i] == "items") { + items = Json::parse(hash[i + 1], nullptr, false); + } + if (hash[i] == "total") { + total = std::stod(hash[i + 1]); + } + } + if (!items.is_array()) { + items = Json::array(); + } + const int product_id = items.empty() ? 1 : items.front().value("product_id", 1); + + auto product = pool_.submit( + [this, product_id] { return http_.get("/products/" + std::to_string(product_id)); }); + auto history = pool_.submit([this, user_id] { + return redis_.string_array({"LRANGE", "bench:history:" + user_id, "0", "4"}); + }); + (void)product.get(); + Json recent = Json::array(); + for (const auto& entry : history.get()) { + Json parsed = Json::parse(entry, nullptr, false); + recent.push_back(parsed.is_discarded() ? Json{{"raw", entry}} : std::move(parsed)); + } + return {{"user_id", user_id}, + {"cart", {{"items", items}, {"item_count", items.size()}, {"estimated_total", total}}}, + {"recent_history", std::move(recent)}, + {"server_type", server_type_}}; + } + + Json checkout(const Json& args) { + const std::string user_id = args.value("user_id", "user-00042"); + Json items = args.value("items", Json::array()); + if (!items.is_array() || items.empty()) { + items = Json::array( + {{{"product_id", 42}, {"quantity", 2}}, {{"product_id", 1337}, {"quantity", 1}}}); + } + const auto now = static_cast(std::time(nullptr)); + const std::string history_key = "bench:history:" + user_id; + const int product_id = items.front().value("product_id", 1); + std::ostringstream rate_key; + rate_key << "bench:ratelimit:user-" << std::setw(5) << std::setfill('0') + << (user_number(user_id) % 100); + const Json order_entry = { + {"order_id", "ORD-" + user_id + "-" + std::to_string(now)}, {"items", items}, {"ts", now}}; + const Json calculate = {{"user_id", user_id}, {"items", items}}; + + auto calculated = pool_.submit( + [this, body = calculate.dump()] { return http_.post("/cart/calculate", body); }); + auto rate = + pool_.submit([this, key = rate_key.str()] { return redis_.integer({"INCR", key}); }); + auto history = pool_.submit([this, history_key, value = order_entry.dump()] { + redis_.discard({"RPUSH", history_key, value}); + }); + auto popularity = pool_.submit([this, product_id] { + redis_.discard({"ZINCRBY", "bench:popular", "1", "product:" + std::to_string(product_id)}); + }); + Json calc = Json::parse(calculated.get()); + const auto rate_count = rate.get(); + history.get(); + popularity.get(); + return {{"order_id", calc.value("order_id", order_entry["order_id"].get())}, + {"user_id", user_id}, + {"total", calc.value("total", 0.0)}, + {"items_count", items.size()}, + {"rate_limit_count", rate_count}, + {"status", "confirmed"}, + {"server_type", server_type_}}; + } + + std::string server_type_; + HttpClient http_; + RedisClient redis_; + ThreadPool pool_; +}; + +Workload::Workload(std::string server_type) : impl_(std::make_unique(std::move(server_type))) {} +Workload::~Workload() = default; + +std::string Workload::invoke(const std::string& tool_name, const std::string& arguments_json) { + return impl_->invoke(tool_name, arguments_json); +} + +std::string Workload::input_schema(const std::string& tool_name) { + if (tool_name == "search_products") { + return R"({"type":"object","properties":{"category":{"type":"string"},"min_price":{"type":"number"},"max_price":{"type":"number"},"limit":{"type":"integer"}}})"; + } + if (tool_name == "get_user_cart") { + return R"({"type":"object","properties":{"user_id":{"type":"string"}}})"; + } + if (tool_name == "checkout") { + return R"({"type":"object","properties":{"user_id":{"type":"string"},"items":{"type":"array","items":{"type":"object","properties":{"product_id":{"type":"integer"},"quantity":{"type":"integer"}},"required":["product_id","quantity"]}}}})"; + } + throw std::invalid_argument("unknown benchmark tool: " + tool_name); +} + +} // namespace mcp_benchmark diff --git a/benchmark/alternatives/common/benchmark_workload.hpp b/benchmark/alternatives/common/benchmark_workload.hpp new file mode 100644 index 0000000..7f983fb --- /dev/null +++ b/benchmark/alternatives/common/benchmark_workload.hpp @@ -0,0 +1,30 @@ +#pragma once + +#include +#include + +namespace mcp_benchmark { + +/// The TM Dev Lab v2 tool workload shared by alternative SDK adapters. +/// +/// The SDK-specific adapter owns MCP registration and transport behavior. This +/// class owns only the Redis + API operations, so every adapter executes the +/// same application work and returns the same JSON text shape. +class Workload { + public: + explicit Workload(std::string server_type); + ~Workload(); + + Workload(const Workload&) = delete; + Workload& operator=(const Workload&) = delete; + + std::string invoke(const std::string& tool_name, const std::string& arguments_json); + + static std::string input_schema(const std::string& tool_name); + + private: + class Impl; + std::unique_ptr impl_; +}; + +} // namespace mcp_benchmark diff --git a/benchmark/alternatives/sources.tsv b/benchmark/alternatives/sources.tsv new file mode 100644 index 0000000..014cedf --- /dev/null +++ b/benchmark/alternatives/sources.tsv @@ -0,0 +1,6 @@ +# id repository commit benchmark_status +hkr04 https://github.com/hkr04/cpp-mcp.git f1117d5286efe6477ddd14322703f034615b6c0e supported +fastmcpp https://github.com/0xeb/fastmcpp.git 9a3ee7125b1db9dc85a71f60661beb40ea1e7993 supported +gopher-mcp https://github.com/GopherSecurity/gopher-mcp.git 63e561107bebad254f854e60c39806a2e6e2468f not-comparable-legacy-http-sse +neumann https://github.com/Neumann-Labs/mcp-cpp.git 57be64093a2241b58d01292d096b266fc6a52ec8 supported +cxxmcp https://github.com/caomengxuan666/cxxmcp.git e48b75dc37372f6f7fc876dc436e34d2909e39af supported diff --git a/benchmark/benchmark_order.py b/benchmark/benchmark_order.py new file mode 100644 index 0000000..d8fc30d --- /dev/null +++ b/benchmark/benchmark_order.py @@ -0,0 +1,57 @@ +#!/usr/bin/env python3 +"""Generate a reproducible, counterbalanced benchmark schedule.""" + +from __future__ import annotations + +import argparse +import json +import random + + +def counterbalanced_orders( + servers: list[str], runs: int, seed: int +) -> list[list[str]]: + base = list(servers) + random.Random(seed).shuffle(base) + if runs == 3 and len(base) == 5: + # Each server occupies three distinct positions, and every pair appears + # in both lead/follow orders over the three-round production profile. + indices = ( + (0, 1, 2, 3, 4), + (1, 2, 3, 4, 0), + (4, 3, 0, 2, 1), + ) + return [[base[index] for index in order] for order in indices] + + return [ + base[offset % len(base) :] + base[: offset % len(base)] + for offset in range(runs) + ] + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--seed", type=int, required=True) + parser.add_argument("--runs", type=int, required=True) + parser.add_argument("--run", type=int) + parser.add_argument("servers", nargs="+") + args = parser.parse_args() + + if args.runs < 1: + parser.error("--runs must be at least 1") + if len(set(args.servers)) != len(args.servers): + parser.error("server names must be unique") + + orders = counterbalanced_orders(args.servers, args.runs, args.seed) + if args.run is None: + print(json.dumps({"seed": args.seed, "orders": orders}, indent=2)) + return 0 + if args.run < 1 or args.run > args.runs: + parser.error("--run must be between 1 and --runs") + for server in orders[args.run - 1]: + print(server) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/capture_container.py b/benchmark/capture_container.py new file mode 100644 index 0000000..88e1368 --- /dev/null +++ b/benchmark/capture_container.py @@ -0,0 +1,437 @@ +#!/usr/bin/env python3 +"""Capture and validate per-run Docker resource and image provenance.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import subprocess +import sys +import tempfile +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Sequence + + +GIBIBYTE = 1_073_741_824 + + +class CaptureError(RuntimeError): + """Raised when Docker metadata is missing or the resource contract is violated.""" + + +def run_command(command: Sequence[str]) -> subprocess.CompletedProcess[str]: + try: + return subprocess.run( + list(command), capture_output=True, text=True, check=False + ) + except OSError as error: + raise CaptureError(f"could not execute {command[0]!r}: {error}") from error + + +def required_command(command: Sequence[str], description: str) -> str: + completed = run_command(command) + if completed.returncode != 0: + detail = completed.stderr.strip() or completed.stdout.strip() or "unknown error" + raise CaptureError(f"could not capture {description}: {detail}") + if not completed.stdout.strip(): + raise CaptureError(f"could not capture {description}: command produced no output") + return completed.stdout + + +def inspect_one( + command: Sequence[str], description: str +) -> tuple[str, dict[str, Any]]: + raw = required_command(command, description) + try: + value = json.loads(raw) + except json.JSONDecodeError as error: + raise CaptureError(f"{description} was not valid JSON") from error + if ( + not isinstance(value, list) + or len(value) != 1 + or not isinstance(value[0], dict) + ): + raise CaptureError(f"expected exactly one {description} record") + return raw, value[0] + + +def container_file(container: str, path: str) -> dict[str, Any]: + completed = run_command(["docker", "exec", container, "cat", path]) + return { + "path": path, + "available": completed.returncode == 0, + "value": completed.stdout.strip() if completed.returncode == 0 else None, + "error": completed.stderr.strip() if completed.returncode != 0 else None, + } + + +def first_available(container: str, paths: Sequence[str]) -> dict[str, Any]: + attempts = [container_file(container, path) for path in paths] + selected = next((attempt for attempt in attempts if attempt["available"]), None) + return {"selected": selected, "attempts": attempts} + + +def detect_cgroup(container: str) -> dict[str, Any]: + version_probe = container_file(container, "/sys/fs/cgroup/cgroup.controllers") + if version_probe["available"]: + return { + "version": 2, + "cpu": first_available(container, ["/sys/fs/cgroup/cpu.max"]), + "memory": first_available(container, ["/sys/fs/cgroup/memory.max"]), + "cpuset_effective": first_available( + container, ["/sys/fs/cgroup/cpuset.cpus.effective"] + ), + "controllers": version_probe, + } + + return { + "version": 1, + "cpu_quota": first_available( + container, + [ + "/sys/fs/cgroup/cpu/cpu.cfs_quota_us", + "/sys/fs/cgroup/cpu,cpuacct/cpu.cfs_quota_us", + ], + ), + "cpu_period": first_available( + container, + [ + "/sys/fs/cgroup/cpu/cpu.cfs_period_us", + "/sys/fs/cgroup/cpu,cpuacct/cpu.cfs_period_us", + ], + ), + "memory": first_available( + container, + [ + "/sys/fs/cgroup/memory/memory.limit_in_bytes", + "/sys/fs/cgroup/memory.limit_in_bytes", + ], + ), + "cpuset_effective": first_available( + container, + [ + "/sys/fs/cgroup/cpuset/cpuset.effective_cpus", + "/sys/fs/cgroup/cpuset/cpuset.cpus", + "/sys/fs/cgroup/cpuset.cpus.effective", + "/sys/fs/cgroup/cpuset.cpus", + ], + ), + "controllers": version_probe, + } + + +def selected_value(record: dict[str, Any]) -> str | None: + selected = record.get("selected") + if not selected: + return None + value = selected.get("value") + return str(value) if value is not None else None + + +def parse_cpu_set(value: str | None) -> set[int] | None: + if value is None or not value.strip(): + return None + cpus: set[int] = set() + try: + for field in value.split(","): + bounds = field.strip().split("-", 1) + start = int(bounds[0]) + end = int(bounds[1]) if len(bounds) == 2 else start + if start < 0 or end < start: + return None + cpus.update(range(start, end + 1)) + except ValueError: + return None + return cpus + + +def validate_resources( + container_inspect: dict[str, Any], + cgroup: dict[str, Any], + expected_cpus: float, + expected_memory_bytes: int, + expected_cpuset: str | None = None, +) -> dict[str, Any]: + host_config = container_inspect.get("HostConfig") or {} + expected_nano_cpus = round(expected_cpus * 1_000_000_000) + checks: dict[str, dict[str, Any]] = { + "host_nano_cpus": { + "expected": expected_nano_cpus, + "actual": host_config.get("NanoCpus"), + }, + "host_memory_bytes": { + "expected": expected_memory_bytes, + "actual": host_config.get("Memory"), + }, + } + checks["host_nano_cpus"]["ok"] = ( + checks["host_nano_cpus"]["actual"] == expected_nano_cpus + ) + checks["host_memory_bytes"]["ok"] = ( + checks["host_memory_bytes"]["actual"] == expected_memory_bytes + ) + + if expected_cpuset is not None: + expected_cpu_set = parse_cpu_set(expected_cpuset) + host_cpu_set = parse_cpu_set(host_config.get("CpusetCpus")) + cgroup_cpu_set = parse_cpu_set( + selected_value(cgroup["cpuset_effective"]) + ) + checks["host_cpuset"] = { + "expected": expected_cpuset, + "actual": host_config.get("CpusetCpus"), + "ok": expected_cpu_set is not None and host_cpu_set == expected_cpu_set, + } + checks["cgroup_cpuset_effective"] = { + "expected": expected_cpuset, + "actual": selected_value(cgroup["cpuset_effective"]), + "ok": expected_cpu_set is not None and cgroup_cpu_set == expected_cpu_set, + } + + if cgroup["version"] == 2: + cpu_value = selected_value(cgroup["cpu"]) + cpu_fields = cpu_value.split() if cpu_value else [] + quota = ( + int(cpu_fields[0]) + if len(cpu_fields) == 2 and cpu_fields[0].isdigit() + else None + ) + period = ( + int(cpu_fields[1]) + if len(cpu_fields) == 2 and cpu_fields[1].isdigit() + else None + ) + else: + quota_value = selected_value(cgroup["cpu_quota"]) + period_value = selected_value(cgroup["cpu_period"]) + quota = ( + int(quota_value) + if quota_value and quota_value.lstrip("-").isdigit() + else None + ) + period = int(period_value) if period_value and period_value.isdigit() else None + + memory_value = selected_value(cgroup["memory"]) + memory_limit = int(memory_value) if memory_value and memory_value.isdigit() else None + checks["cgroup_cpu_quota"] = { + "expected_cpus": expected_cpus, + "quota": quota, + "period": period, + "ok": ( + quota is not None + and period is not None + and period > 0 + and quota == round(expected_cpus * period) + ), + } + checks["cgroup_memory_bytes"] = { + "expected": expected_memory_bytes, + "actual": memory_limit, + "ok": memory_limit == expected_memory_bytes, + } + errors = [name for name, check in checks.items() if not check["ok"]] + return {"ok": not errors, "checks": checks, "failed_checks": errors} + + +def configured_executable(container_inspect: dict[str, Any]) -> str | None: + config = container_inspect.get("Config") or {} + entrypoint = config.get("Entrypoint") or [] + command = config.get("Cmd") or [] + if isinstance(entrypoint, str): + entrypoint = [entrypoint] + if isinstance(command, str): + command = [command] + argv = [*entrypoint, *command] + return str(argv[0]) if argv else None + + +def executable_provenance(container: str, configured: str | None) -> dict[str, Any]: + record: dict[str, Any] = { + "configured": configured, + "resolved_path": None, + "sha256": None, + } + if not configured: + record["error"] = "container has no configured command" + return record + + if configured.startswith("/"): + resolved = configured + else: + resolver = run_command( + [ + "docker", + "exec", + container, + "sh", + "-c", + 'candidate=$(command -v "$1") || exit 1; case "$candidate" in /*) printf "%s\\n" "$candidate" ;; *) printf "%s/%s\\n" "$PWD" "$candidate" ;; esac', + "resolve", + configured, + ] + ) + if resolver.returncode != 0 or not resolver.stdout.strip(): + record["error"] = ( + resolver.stderr.strip() + or "configured executable was not resolvable" + ) + return record + resolved = resolver.stdout.strip().splitlines()[0] + record["resolved_path"] = resolved + + with tempfile.TemporaryDirectory(prefix="mcp-benchmark-executable-") as temp_dir: + destination = Path(temp_dir) / "executable" + copier = run_command( + ["docker", "cp", "-L", f"{container}:{resolved}", str(destination)] + ) + if copier.returncode != 0 or not destination.is_file(): + record["error"] = copier.stderr.strip() or "docker cp produced no file" + return record + record["sha256"] = hashlib.sha256(destination.read_bytes()).hexdigest() + return record + + +def positive_integer(value: str) -> int: + parsed = int(value) + if parsed <= 0: + raise argparse.ArgumentTypeError("must be greater than zero") + return parsed + + +def positive_number(value: str) -> float: + parsed = float(value) + if parsed <= 0: + raise argparse.ArgumentTypeError("must be greater than zero") + return parsed + + +def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description=( + "Write container_inspect_runN.json, image_inspect_runN.json, and " + "cgroup_runN.json, then enforce the benchmark resource contract." + ) + ) + parser.add_argument( + "--container", required=True, help="running container name or ID" + ) + parser.add_argument("--run", required=True, type=positive_integer, dest="run_number") + parser.add_argument("--output-dir", required=True, type=Path) + parser.add_argument("--expected-cpus", type=positive_number, default=2.0) + parser.add_argument( + "--expected-memory-bytes", type=positive_integer, default=2 * GIBIBYTE + ) + parser.add_argument( + "--expected-cpuset", + help=( + "CPU list/ranges that Docker and the effective cgroup must exactly match" + ), + ) + parser.add_argument( + "--skip-executable", + action="store_true", + help="skip docker cp hashing when a pinned image digest is sufficient", + ) + return parser.parse_args(argv) + + +def capture(args: argparse.Namespace) -> dict[str, Any]: + raw_container, container_inspect = inspect_one( + ["docker", "inspect", args.container], + f"container inspect for {args.container!r}", + ) + image_reference = container_inspect.get("Image") + if not image_reference: + raise CaptureError("container inspect did not include an image ID") + raw_image, image_inspect = inspect_one( + ["docker", "image", "inspect", str(image_reference)], + f"image inspect for {image_reference!r}", + ) + + args.output_dir.mkdir(parents=True, exist_ok=True) + suffix = f"run{args.run_number}" + (args.output_dir / f"container_inspect_{suffix}.json").write_text( + raw_container.rstrip() + "\n", encoding="utf-8" + ) + (args.output_dir / f"image_inspect_{suffix}.json").write_text( + raw_image.rstrip() + "\n", encoding="utf-8" + ) + + cgroup = detect_cgroup(args.container) + validation = validate_resources( + container_inspect, + cgroup, + args.expected_cpus, + args.expected_memory_bytes, + args.expected_cpuset, + ) + executable = ( + {"skipped": True, "reason": "pinned image provenance is sufficient"} + if args.skip_executable + else executable_provenance( + args.container, configured_executable(container_inspect) + ) + ) + if not args.skip_executable and not executable.get("sha256"): + raise CaptureError( + "could not capture configured executable: " + f"{executable.get('error', 'missing SHA-256')}" + ) + + summary = { + "schema_version": 1, + "captured_at_utc": datetime.now(timezone.utc).isoformat(), + "container": { + "requested": args.container, + "id": container_inspect.get("Id"), + "name": container_inspect.get("Name"), + "state": (container_inspect.get("State") or {}).get("Status"), + "host_config": { + "nano_cpus": (container_inspect.get("HostConfig") or {}).get( + "NanoCpus" + ), + "memory_bytes": (container_inspect.get("HostConfig") or {}).get( + "Memory" + ), + "cpuset_cpus": (container_inspect.get("HostConfig") or {}).get( + "CpusetCpus" + ), + }, + }, + "image": { + "id": image_inspect.get("Id"), + "repo_tags": image_inspect.get("RepoTags") or [], + "repo_digests": image_inspect.get("RepoDigests") or [], + }, + "executable": executable, + "cgroup": cgroup, + "cpuset_effective": selected_value(cgroup["cpuset_effective"]), + "validation": validation, + } + (args.output_dir / f"cgroup_{suffix}.json").write_text( + json.dumps(summary, indent=2) + "\n", encoding="utf-8" + ) + return summary + + +def main(argv: Sequence[str] | None = None) -> int: + args = parse_args(argv) + try: + summary = capture(args) + except CaptureError as error: + print(f"capture_container.py: {error}", file=sys.stderr) + return 2 + if not summary["validation"]["ok"]: + failed = ", ".join(summary["validation"]["failed_checks"]) + print( + f"capture_container.py: resource validation failed: {failed}", + file=sys.stderr, + ) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/capture_environment.py b/benchmark/capture_environment.py new file mode 100644 index 0000000..545fc68 --- /dev/null +++ b/benchmark/capture_environment.py @@ -0,0 +1,500 @@ +#!/usr/bin/env python3 +"""Capture reproducible host, source, dependency, and profile provenance.""" + +from __future__ import annotations + +import argparse +import csv +import hashlib +import io +import json +import os +import platform +import subprocess +import sys +import tarfile +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Iterable, Sequence + + +EXCLUDED_BENCHMARK_PARTS = frozenset( + {"__pycache__", "alternative-sdks", "benchmark-mcp-servers-v2", "results"} +) + + +class CaptureError(RuntimeError): + """Raised when required benchmark provenance cannot be captured.""" + + +def run_command(command: Sequence[str], cwd: Path) -> dict[str, Any]: + """Run a command without a shell and retain audit-friendly output.""" + try: + completed = subprocess.run( + list(command), cwd=cwd, capture_output=True, text=True, check=False + ) + except OSError as error: + return { + "command": list(command), + "exit_code": None, + "stdout": "", + "stderr": str(error), + } + return { + "command": list(command), + "exit_code": completed.returncode, + "stdout": completed.stdout.strip(), + "stderr": completed.stderr.strip(), + } + + +def required_stdout(result: dict[str, Any], description: str) -> str: + if result["exit_code"] != 0 or not result["stdout"]: + detail = result["stderr"] or "command produced no output" + raise CaptureError(f"could not capture {description}: {detail}") + return str(result["stdout"]) + + +def read_optional(path: Path) -> str | None: + try: + return path.read_text(encoding="utf-8", errors="replace").strip() + except OSError: + return None + + +def parse_os_release(path: Path = Path("/etc/os-release")) -> dict[str, str]: + values: dict[str, str] = {} + content = read_optional(path) + if content is None: + return values + for line in content.splitlines(): + if "=" not in line or line.startswith("#"): + continue + key, value = line.split("=", 1) + values[key] = value.strip().strip('"') + return values + + +def parse_cpuinfo(path: Path = Path("/proc/cpuinfo")) -> dict[str, Any]: + content = read_optional(path) + if content is None: + return { + "model": None, + "physical_package_count": None, + "physical_core_count": None, + } + + records: list[dict[str, str]] = [] + current: dict[str, str] = {} + for line in content.splitlines(): + if not line.strip(): + if current: + records.append(current) + current = {} + continue + if ":" not in line: + continue + key, value = line.split(":", 1) + current[key.strip()] = value.strip() + if current: + records.append(current) + + model = next( + ( + record.get("model name") + or record.get("Hardware") + or record.get("Processor") + for record in records + if record.get("model name") + or record.get("Hardware") + or record.get("Processor") + ), + None, + ) + package_ids = { + record["physical id"] for record in records if "physical id" in record + } + core_ids = { + (record["physical id"], record["core id"]) + for record in records + if "physical id" in record and "core id" in record + } + return { + "model": model, + "physical_package_count": len(package_ids) if package_ids else None, + "physical_core_count": len(core_ids) if core_ids else None, + } + + +def parse_memory_total_bytes(path: Path = Path("/proc/meminfo")) -> int | None: + content = read_optional(path) + if content is None: + return None + for line in content.splitlines(): + if not line.startswith("MemTotal:"): + continue + fields = line.split() + if len(fields) >= 2 and fields[1].isdigit(): + return int(fields[1]) * 1024 + return None + + +def compress_cpu_list(cpus: Iterable[int]) -> str: + values = sorted(set(cpus)) + if not values: + return "" + ranges: list[str] = [] + start = previous = values[0] + for value in values[1:]: + if value == previous + 1: + previous = value + continue + ranges.append(str(start) if start == previous else f"{start}-{previous}") + start = previous = value + ranges.append(str(start) if start == previous else f"{start}-{previous}") + return ",".join(ranges) + + +def cpu_affinity() -> list[int] | None: + try: + return sorted(os.sched_getaffinity(0)) + except (AttributeError, OSError): + return None + + +def cpu_governors(cpus: Iterable[int] | None) -> dict[str, Any]: + cpu_values = list(cpus) if cpus is not None else list(range(os.cpu_count() or 0)) + by_cpu: dict[str, str] = {} + for cpu in cpu_values: + value = read_optional( + Path(f"/sys/devices/system/cpu/cpu{cpu}/cpufreq/scaling_governor") + ) + if value: + by_cpu[str(cpu)] = value + return {"values": sorted(set(by_cpu.values())), "by_cpu": by_cpu} + + +def benchmark_source_files(project_dir: Path) -> Iterable[Path]: + """Yield core sources and all local harness inputs, excluding fetched/runtime data.""" + for relative in ("CMakeLists.txt", "VERSION", "LICENSE"): + path = project_dir / relative + if path.is_file(): + yield path + + for relative in ("cmake", "include", "src"): + root = project_dir / relative + if not root.is_dir(): + continue + for path in sorted(root.rglob("*")): + if path.is_file() and "__pycache__" not in path.parts: + yield path + + benchmark_root = project_dir / "benchmark" + if not benchmark_root.is_dir(): + return + for path in sorted(benchmark_root.rglob("*")): + if not path.is_file() or path.suffix == ".pyc": + continue + relative_parts = path.relative_to(benchmark_root).parts + if any(part in EXCLUDED_BENCHMARK_PARTS for part in relative_parts): + continue + yield path + + +def tree_digest(project_dir: Path) -> tuple[str, list[str]]: + digest = hashlib.sha256() + relative_paths: list[str] = [] + for path in benchmark_source_files(project_dir): + relative = path.relative_to(project_dir).as_posix() + relative_bytes = relative.encode("utf-8") + content = path.read_bytes() + digest.update(len(relative_bytes).to_bytes(8, "big")) + digest.update(relative_bytes) + digest.update(len(content).to_bytes(8, "big")) + digest.update(content) + relative_paths.append(relative) + return digest.hexdigest(), relative_paths + + +def create_source_snapshot( + project_dir: Path, relative_paths: Sequence[str], output: Path +) -> str: + """Write a deterministic archive of every local source input used by the run.""" + output.parent.mkdir(parents=True, exist_ok=True) + with tarfile.open(output, "w", format=tarfile.PAX_FORMAT) as archive: + for relative in relative_paths: + path = project_dir / relative + content = path.read_bytes() + info = tarfile.TarInfo(relative) + info.size = len(content) + info.mode = path.stat().st_mode & 0o777 + info.mtime = 0 + info.uid = 0 + info.gid = 0 + info.uname = "" + info.gname = "" + archive.addfile(info, io.BytesIO(content)) + return hashlib.sha256(output.read_bytes()).hexdigest() + + +def verify_source_digest(environment_path: Path, project_dir: Path) -> int: + try: + environment = json.loads(environment_path.read_text(encoding="utf-8")) + expected = environment["source"]["tree_sha256"] + except (OSError, json.JSONDecodeError, KeyError, TypeError) as error: + print(f"capture_environment.py: invalid environment record: {error}", file=sys.stderr) + return 2 + actual, _ = tree_digest(project_dir.resolve()) + if actual != expected: + print( + "capture_environment.py: local benchmark sources changed during the run", + file=sys.stderr, + ) + return 1 + return 0 + + +def git_repository( + path: Path, expected_commit: str | None = None +) -> dict[str, Any]: + if not path.is_dir(): + return { + "path": str(path), + "available": False, + "expected_commit": expected_commit, + "error": "directory does not exist", + } + + head_result = run_command(["git", "rev-parse", "HEAD"], path) + status_result = run_command( + ["git", "status", "--porcelain=v1", "--untracked-files=all"], path + ) + origin_result = run_command(["git", "remote", "get-url", "origin"], path) + head = head_result["stdout"] if head_result["exit_code"] == 0 else None + status = status_result["stdout"] if status_result["exit_code"] == 0 else None + return { + "path": str(path.resolve()), + "available": head is not None, + "origin_url": origin_result["stdout"] if origin_result["exit_code"] == 0 else None, + "head": head, + "expected_commit": expected_commit, + "commit_matches": expected_commit is None or head == expected_commit, + "clean": status == "" if status is not None else None, + "status_porcelain": status, + "errors": { + "head": head_result["stderr"] if head_result["exit_code"] != 0 else None, + "status": status_result["stderr"] if status_result["exit_code"] != 0 else None, + "origin": origin_result["stderr"] if origin_result["exit_code"] != 0 else None, + }, + } + + +def alternative_repositories(manifest: Path, root: Path) -> list[dict[str, Any]]: + try: + manifest_file = manifest.open(encoding="utf-8", newline="") + except OSError as error: + raise CaptureError(f"could not open alternative manifest {manifest}: {error}") from error + + with manifest_file: + reader = csv.DictReader(manifest_file, delimiter="\t") + if reader.fieldnames: + reader.fieldnames = [field.removeprefix("#").strip() for field in reader.fieldnames] + required_fields = {"id", "repository", "commit", "benchmark_status"} + if not reader.fieldnames or not required_fields.issubset(reader.fieldnames): + missing = sorted(required_fields - set(reader.fieldnames or [])) + raise CaptureError(f"alternative manifest {manifest} is missing fields: {missing}") + + entries: list[dict[str, Any]] = [] + for row in reader: + identifier = row["id"].strip() + repository = git_repository(root / identifier, row["commit"].strip()) + repository.update( + { + "id": identifier, + "manifest_url": row["repository"].strip(), + "origin_matches_manifest": repository.get("origin_url") + == row["repository"].strip(), + "benchmark_status": row["benchmark_status"].strip(), + } + ) + entries.append(repository) + return entries + + +def inspect_k6_image(reference: str, cwd: Path) -> dict[str, Any]: + result = run_command(["docker", "image", "inspect", reference], cwd) + stdout = required_stdout(result, f"k6 image {reference!r}") + try: + inspected = json.loads(stdout) + except json.JSONDecodeError as error: + raise CaptureError(f"docker returned invalid JSON for k6 image {reference!r}") from error + if not isinstance(inspected, list) or len(inspected) != 1: + raise CaptureError(f"expected one image inspect record for k6 image {reference!r}") + image = inspected[0] + return { + "reference": reference, + "id": image.get("Id"), + "repo_tags": image.get("RepoTags") or [], + "repo_digests": image.get("RepoDigests") or [], + "inspect": image, + } + + +def positive_integer(value: str) -> int: + parsed = int(value) + if parsed <= 0: + raise argparse.ArgumentTypeError("must be greater than zero") + return parsed + + +def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Capture benchmark profile, host, source, and dependency provenance." + ) + parser.add_argument("output", type=Path, help="environment JSON output path") + parser.add_argument("--project-dir", type=Path, required=True) + parser.add_argument("--upstream-dir", type=Path, required=True) + parser.add_argument("--upstream-url") + parser.add_argument("--upstream-commit", required=True) + parser.add_argument("--alternative-root", type=Path, required=True) + parser.add_argument("--sources-manifest", type=Path) + parser.add_argument("--protocol-version", required=True) + parser.add_argument("--eligibility-contract", required=True) + parser.add_argument("--k6-image", required=True) + parser.add_argument("--order-seed", required=True, type=int) + parser.add_argument("--runs", required=True, type=positive_integer) + parser.add_argument("--servers", nargs="+", required=True) + parser.add_argument("--vus", required=True, type=positive_integer) + parser.add_argument( + "--measurement-duration", + "--measure-duration", + dest="measurement_duration", + required=True, + ) + parser.add_argument("--warmup-duration", required=True) + parser.add_argument("--ramp-duration", required=True) + return parser.parse_args(argv) + + +def build_metadata(args: argparse.Namespace) -> dict[str, Any]: + project_dir = args.project_dir.resolve() + upstream_dir = args.upstream_dir.resolve() + alternative_root = args.alternative_root.resolve() + manifest = ( + args.sources_manifest.resolve() + if args.sources_manifest + else project_dir / "benchmark" / "alternatives" / "sources.tsv" + ) + + source_sha256, source_paths = tree_digest(project_dir) + uname = platform.uname() + affinity = cpu_affinity() + cpu = parse_cpuinfo() + + upstream = git_repository(upstream_dir, args.upstream_commit) + upstream["expected_url"] = args.upstream_url + upstream["origin_matches_expected"] = ( + args.upstream_url is None or upstream.get("origin_url") == args.upstream_url + ) + alternatives = alternative_repositories(manifest, alternative_root) + + docker_version = run_command( + ["docker", "version", "--format", "{{json .}}"], project_dir + ) + if docker_version["exit_code"] == 0 and docker_version["stdout"]: + try: + docker_version["parsed"] = json.loads(docker_version["stdout"]) + except json.JSONDecodeError: + docker_version["parsed"] = None + + return { + "schema_version": 2, + "captured_at_utc": datetime.now(timezone.utc).isoformat(), + "profile": { + "servers": args.servers, + "runs": args.runs, + "order_seed": args.order_seed, + "requested_protocol_version": args.protocol_version, + "eligibility_contract": args.eligibility_contract, + "virtual_users": args.vus, + "ramp_duration": args.ramp_duration, + "warmup_duration": args.warmup_duration, + "measurement_duration": args.measurement_duration, + "k6_image_reference": args.k6_image, + }, + "source": { + "tree_sha256": source_sha256, + "file_count": len(source_paths), + "files": source_paths, + "project_repository": git_repository(project_dir), + }, + "dependencies": { + "upstream": upstream, + "alternative_manifest": str(manifest), + "alternatives": alternatives, + }, + "host": { + "os": { + "system": uname.system, + "release": uname.release, + "version": uname.version, + "distribution": parse_os_release(), + }, + "kernel": uname.release, + "architecture": uname.machine, + "cpu": { + **cpu, + "logical_cpu_count": os.cpu_count(), + "affinity": affinity, + "affinity_list": compress_cpu_list(affinity or []), + "governors": cpu_governors(affinity), + "lscpu": run_command(["lscpu", "--json"], project_dir), + }, + "memory_total_bytes": parse_memory_total_bytes(), + }, + "tools": { + "docker": docker_version, + "docker_compose": run_command(["docker", "compose", "version"], project_dir), + "python": { + "version": platform.python_version(), + "executable": sys.executable, + "command": run_command([sys.executable, "--version"], project_dir), + }, + "git": run_command(["git", "--version"], project_dir), + }, + "k6_image": inspect_k6_image(args.k6_image, project_dir), + } + + +def main(argv: Sequence[str] | None = None) -> int: + effective_argv = list(argv) if argv is not None else sys.argv[1:] + if effective_argv and effective_argv[0] == "--verify-source": + verifier = argparse.ArgumentParser() + verifier.add_argument("--verify-source", type=Path, required=True) + verifier.add_argument("--project-dir", type=Path, required=True) + verify_args = verifier.parse_args(effective_argv) + return verify_source_digest(verify_args.verify_source, verify_args.project_dir) + + args = parse_args(effective_argv) + try: + metadata = build_metadata(args) + except CaptureError as error: + print(f"capture_environment.py: {error}", file=sys.stderr) + return 2 + + args.output.parent.mkdir(parents=True, exist_ok=True) + snapshot_path = args.output.parent / "source_snapshot.tar" + snapshot_sha256 = create_source_snapshot( + args.project_dir.resolve(), metadata["source"]["files"], snapshot_path + ) + metadata["source"]["snapshot"] = { + "path": snapshot_path.name, + "sha256": snapshot_sha256, + "format": "deterministic POSIX tar", + } + args.output.write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/collect_stats.py b/benchmark/collect_stats.py index dc0f990..dd2f227 100755 --- a/benchmark/collect_stats.py +++ b/benchmark/collect_stats.py @@ -1,18 +1,23 @@ #!/usr/bin/env python3 +"""Collect auditable resource samples for a container or the benchmark host.""" + +from __future__ import annotations import json +import os import re import signal import subprocess import sys import time from datetime import datetime, timezone +from decimal import Decimal from pathlib import Path +from typing import Any RUNNING = True -SAMPLES = [] -OUTPUT_PATH: Path | None = None +STOP_SIGNAL: int | None = None SIZE_UNITS = { @@ -26,6 +31,19 @@ "gib": 1024**3, "tib": 1024**4, } +ANSI_CONTROL = re.compile(r"\x1b\[[0-9;?]*[A-Za-z]") +SIZE_VALUE = re.compile( + r"^((?:[0-9]+(?:\.[0-9]*)?|\.[0-9]+)(?:e[+-]?[0-9]+)?)\s*([a-z]+)$" +) +HOST_NETWORK_POLICY = "baseline-cohort-retire-v1" + + +class CollectionStopped(Exception): + """Internal signal that a streaming collector was stopped intentionally.""" + + +def utc_now() -> str: + return datetime.now(timezone.utc).isoformat() def parse_size_to_bytes(value: str) -> int: @@ -33,11 +51,11 @@ def parse_size_to_bytes(value: str) -> int: if text in {"", "0", "0b", "--"}: return 0 - match = re.match(r"^([0-9]*\.?[0-9]+)\s*([a-z]+)$", text) + match = SIZE_VALUE.match(text) if not match: raise ValueError(f"Unable to parse size value: {value!r}") - number = float(match.group(1)) + number = Decimal(match.group(1)) unit = match.group(2) if unit not in SIZE_UNITS: raise ValueError(f"Unknown size unit in value: {value!r}") @@ -51,53 +69,83 @@ def parse_cpu_percent(value: str) -> float: def parse_mem_usage(value: str) -> tuple[int, int]: - parts = [p.strip() for p in value.split("/")] + parts = [part.strip() for part in value.split("/")] if len(parts) != 2: raise ValueError(f"Unexpected memory usage format: {value!r}") return parse_size_to_bytes(parts[0]), parse_size_to_bytes(parts[1]) def parse_net_io(value: str) -> tuple[int, int]: - parts = [p.strip() for p in value.split("/")] + parts = [part.strip() for part in value.split("/")] if len(parts) != 2: raise ValueError(f"Unexpected net I/O format: {value!r}") return parse_size_to_bytes(parts[0]), parse_size_to_bytes(parts[1]) -def write_samples() -> None: - global OUTPUT_PATH - if OUTPUT_PATH is None: - return - - OUTPUT_PATH.parent.mkdir(parents=True, exist_ok=True) - with OUTPUT_PATH.open("w", encoding="utf-8") as f: - json.dump(SAMPLES, f, indent=2) +def audit_path_for(output_path: Path) -> Path: + """Return the deterministic sidecar path without changing the sample format.""" + return output_path.with_name(f"{output_path.stem}.audit.json") + + +def readiness_path_for(output_path: Path) -> Path: + """Return the marker written only after the first valid baseline sample.""" + return output_path.with_name(f"{output_path.stem}.ready.json") + + +def write_json(path: Path, value: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", encoding="utf-8") as stream: + json.dump(value, stream, indent=2) + stream.write("\n") + + +def write_initialization_failure( + output_path: Path, + target: str, + target_type: str, + interval_seconds: float, + error: Exception, +) -> bool: + """Best-effort failure artifacts for a collector that could not start.""" + timestamp = utc_now() + audit = { + "schema_version": 1, + "status": "failed", + "target": target, + "target_type": target_type, + "output_path": str(output_path), + "audit_path": str(audit_path_for(output_path)), + "readiness_path": str(readiness_path_for(output_path)), + "interval_seconds": interval_seconds, + "started_at": timestamp, + "finished_at": timestamp, + "elapsed_seconds": 0.0, + "termination": {"reason": "initialization_failure", "signal": None}, + "attempt_count": 0, + "sample_count": 0, + "failure_count": 1, + "failures": [{"timestamp": timestamp, "error": str(error)}], + "fatal_error": str(error), + } + try: + write_json(output_path, []) + write_json(audit_path_for(output_path), audit) + except Exception as write_error: + print( + f"Failed to write collector initialization artifacts: {write_error}", + file=sys.stderr, + ) + return False + return True -def handle_signal(_signum: int, _frame) -> None: - global RUNNING +def handle_signal(signum: int, _frame: object) -> None: + global RUNNING, STOP_SIGNAL + STOP_SIGNAL = signum RUNNING = False -def collect_once(container_name: str) -> dict: - format_str = "{{.CPUPerc}}|{{.MemUsage}}|{{.NetIO}}" - cmd = [ - "docker", - "stats", - "--no-stream", - "--format", - format_str, - container_name, - ] - proc = subprocess.run(cmd, capture_output=True, text=True, check=False) - - if proc.returncode != 0: - raise RuntimeError(proc.stderr.strip() or proc.stdout.strip() or "docker stats failed") - - line = proc.stdout.strip() - if not line: - raise RuntimeError("docker stats returned empty output") - +def parse_container_stats_line(line: str) -> dict[str, int | float | str]: parts = line.split("|") if len(parts) != 3: raise RuntimeError(f"Unexpected docker stats format: {line!r}") @@ -105,9 +153,8 @@ def collect_once(container_name: str) -> dict: cpu_percent = parse_cpu_percent(parts[0]) mem_usage_bytes, mem_limit_bytes = parse_mem_usage(parts[1]) net_io_rx, net_io_tx = parse_net_io(parts[2]) - return { - "timestamp": datetime.now(timezone.utc).isoformat(), + "timestamp": utc_now(), "cpu_percent": cpu_percent, "mem_usage_bytes": mem_usage_bytes, "mem_limit_bytes": mem_limit_bytes, @@ -116,18 +163,281 @@ def collect_once(container_name: str) -> dict: } +def collect_container_once(container_name: str) -> dict[str, int | float | str]: + """Collect one container sample; retained as a deterministic test seam.""" + result = subprocess.run( + [ + "docker", + "stats", + "--no-stream", + "--format", + "{{.CPUPerc}}|{{.MemUsage}}|{{.NetIO}}", + container_name, + ], + check=True, + capture_output=True, + text=True, + ) + lines = [ + ANSI_CONTROL.sub("", line).strip() + for line in result.stdout.splitlines() + if ANSI_CONTROL.sub("", line).strip() + ] + if len(lines) != 1: + raise RuntimeError( + f"docker stats returned {len(lines)} rows for {container_name!r}" + ) + return parse_container_stats_line(lines[0]) + + +_DEFAULT_COLLECT_CONTAINER_ONCE = collect_container_once + + +class ContainerStatsStream: + """Read a persistent Docker stats stream without repeated startup delays.""" + + def __init__(self, container_name: str) -> None: + self.container_name = container_name + self.process = subprocess.Popen( + [ + "docker", + "stats", + "--format", + "{{.CPUPerc}}|{{.MemUsage}}|{{.NetIO}}", + container_name, + ], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + bufsize=1, + ) + if self.process.stdout is None or self.process.stderr is None: + self.close() + raise RuntimeError("docker stats stream did not expose stdout and stderr") + + def collect(self) -> dict[str, int | float | str]: + assert self.process.stdout is not None + while RUNNING: + raw_line = self.process.stdout.readline() + if not RUNNING: + raise CollectionStopped + if raw_line == "": + return_code = self.process.poll() + detail = self._stderr() + raise RuntimeError( + detail + or f"docker stats stream exited unexpectedly with code {return_code}" + ) + line = ANSI_CONTROL.sub("", raw_line).strip() + if not line: + continue + return parse_container_stats_line(line) + raise CollectionStopped + + def _stderr(self) -> str: + if self.process.stderr is None: + return "" + return self.process.stderr.read().strip() + + def close(self) -> None: + if self.process.poll() is not None: + return + self.process.terminate() + try: + self.process.wait(timeout=5) + except subprocess.TimeoutExpired: + self.process.kill() + self.process.wait(timeout=5) + + +def read_cpu_snapshot(proc_root: Path, affinity: frozenset[int]) -> tuple[int, int]: + """Return aggregate (total, busy) scheduler ticks for affinity-visible CPUs.""" + totals = 0 + busy = 0 + observed: set[int] = set() + with (proc_root / "stat").open(encoding="utf-8") as stream: + for line in stream: + match = re.match(r"^cpu(\d+)\s+(.+)$", line.rstrip()) + if not match: + continue + cpu = int(match.group(1)) + if cpu not in affinity: + continue + values = [int(value) for value in match.group(2).split()] + if len(values) < 5: + raise RuntimeError(f"Malformed /proc/stat row for cpu{cpu}") + # Linux accounts guest time inside user/nice, so only the first eight + # counters participate in the conventional non-double-counted total. + cpu_total = sum(values[:8]) + cpu_idle = values[3] + values[4] + totals += cpu_total + busy += cpu_total - cpu_idle + observed.add(cpu) + missing = affinity - observed + if missing: + raise RuntimeError(f"/proc/stat is missing affinity CPUs: {sorted(missing)}") + return totals, busy + + +def read_host_memory(proc_root: Path) -> tuple[int, int]: + fields: dict[str, int] = {} + with (proc_root / "meminfo").open(encoding="utf-8") as stream: + for line in stream: + key, separator, value = line.partition(":") + if not separator: + continue + parts = value.split() + if parts: + fields[key] = int(parts[0]) * 1024 + try: + total = fields["MemTotal"] + available = fields["MemAvailable"] + except KeyError as error: + raise RuntimeError(f"Missing {error.args[0]} in /proc/meminfo") from error + return total - available, total + + +def read_host_network(proc_root: Path) -> dict[str, tuple[int, int]]: + counters: dict[str, tuple[int, int]] = {} + with (proc_root / "net" / "dev").open(encoding="utf-8") as stream: + for line in stream: + if ":" not in line: + continue + interface, values = line.split(":", 1) + interface = interface.strip() + fields = values.split() + if len(fields) < 16: + raise RuntimeError(f"Malformed /proc/net/dev row: {line.rstrip()!r}") + if not interface or interface in counters: + raise RuntimeError( + f"Invalid interface in /proc/net/dev row: {line.rstrip()!r}" + ) + counters[interface] = (int(fields[0]), int(fields[8])) + if not counters: + raise RuntimeError("/proc/net/dev contains no network interfaces") + return counters + + +class HostStatsCollector: + """Measure the host resources available to this affinity-constrained process.""" + + def __init__( + self, + proc_root: Path = Path("/proc"), + affinity: frozenset[int] | None = None, + ) -> None: + self.proc_root = proc_root + self.affinity = ( + frozenset(os.sched_getaffinity(0)) if affinity is None else affinity + ) + if not self.affinity: + raise RuntimeError("collector process has an empty CPU affinity") + self._previous_total, self._previous_busy = read_cpu_snapshot( + self.proc_root, self.affinity + ) + network = read_host_network(self.proc_root) + self.network_interfaces = frozenset(network) + self.active_network_interfaces = set(self.network_interfaces) + self.retired_network_interfaces: set[str] = set() + self.ignored_new_network_interfaces: set[str] = set() + self._previous_network = { + interface: network[interface] for interface in self.network_interfaces + } + + def audit_metadata(self) -> dict[str, Any]: + """Describe the fixed host scope and any observed interface retirement.""" + return { + "cpu_affinity": sorted(self.affinity), + "cpu_capacity_cores": len(self.affinity), + "network_interfaces": sorted(self.network_interfaces), + "active_network_interfaces": sorted(self.active_network_interfaces), + "retired_network_interfaces": sorted(self.retired_network_interfaces), + "ignored_new_network_interfaces": sorted( + self.ignored_new_network_interfaces + ), + "network_interface_policy": HOST_NETWORK_POLICY, + "network_interface_accounting": ( + "baseline cohort; freeze last counters on retirement; " + "ignore later additions" + ), + } + + def collect(self) -> dict[str, int | float | str]: + total, busy = read_cpu_snapshot(self.proc_root, self.affinity) + total_delta = total - self._previous_total + busy_delta = busy - self._previous_busy + self._previous_total, self._previous_busy = total, busy + if total_delta <= 0 or busy_delta < 0: + raise RuntimeError("host CPU counters did not advance monotonically") + + # Match docker-stats semantics: one fully busy core is 100%, so the + # capacity is affinity-core-count * 100%. + cpu_percent = busy_delta / total_delta * len(self.affinity) * 100.0 + mem_usage_bytes, mem_limit_bytes = read_host_memory(self.proc_root) + network = read_host_network(self.proc_root) + self.ignored_new_network_interfaces.update( + network.keys() - self.network_interfaces + ) + missing_interfaces = self.active_network_interfaces - network.keys() + if missing_interfaces: + # Docker can remove a short-lived host-side interface after the + # collector has taken its baseline. Freeze that interface at its + # last observed counters so the aggregate remains monotonic and the + # bytes observed before retirement remain represented. Interfaces + # created after the baseline are intentionally never adopted. + self.active_network_interfaces.difference_update(missing_interfaces) + self.retired_network_interfaces.update(missing_interfaces) + if not self.active_network_interfaces: + raise RuntimeError( + "all baseline host network interfaces disappeared during collection" + ) + for interface in self.active_network_interfaces: + previous_rx, previous_tx = self._previous_network[interface] + current_rx, current_tx = network[interface] + if current_rx < previous_rx or current_tx < previous_tx: + raise RuntimeError( + f"host network counters reset for interface {interface!r}" + ) + self._previous_network[interface] = (current_rx, current_tx) + net_io_rx = sum( + self._previous_network[interface][0] + for interface in self.network_interfaces + ) + net_io_tx = sum( + self._previous_network[interface][1] + for interface in self.network_interfaces + ) + return { + "timestamp": utc_now(), + "cpu_percent": cpu_percent, + "mem_usage_bytes": mem_usage_bytes, + "mem_limit_bytes": mem_limit_bytes, + "net_io_rx": net_io_rx, + "net_io_tx": net_io_tx, + } + + +def wait_until(deadline: float) -> None: + while RUNNING: + remaining = deadline - time.monotonic() + if remaining <= 0: + return + time.sleep(min(0.2, remaining)) + + def main() -> int: - global OUTPUT_PATH + global RUNNING, STOP_SIGNAL if len(sys.argv) != 4: print( - "Usage: collect_stats.py ", + "Usage: collect_stats.py " + " ", file=sys.stderr, ) return 1 - container_name = sys.argv[1] - OUTPUT_PATH = Path(sys.argv[2]) + target = sys.argv[1] + output_path = Path(sys.argv[2]) try: interval_seconds = float(sys.argv[3]) if interval_seconds <= 0: @@ -136,34 +446,148 @@ def main() -> int: print(f"Invalid interval_seconds: {exc}", file=sys.stderr) return 1 + RUNNING = True + STOP_SIGNAL = None signal.signal(signal.SIGTERM, handle_signal) signal.signal(signal.SIGINT, handle_signal) - exit_code = 0 + target_type = "host" if target == "@host" else "container" + host_collector: HostStatsCollector | None = None + container_stream: ContainerStatsStream | None = None + injected_single_sample = collect_container_once is not _DEFAULT_COLLECT_CONTAINER_ONCE + try: + if target_type == "host": + host_collector = HostStatsCollector() + elif not injected_single_sample: + container_stream = ContainerStatsStream(target) + except Exception as exc: + print(f"Failed to initialize resource collector: {exc}", file=sys.stderr) + write_initialization_failure( + output_path, target, target_type, interval_seconds, exc + ) + return 1 + + started_at = utc_now() + started_monotonic = time.monotonic() + samples: list[dict[str, int | float | str]] = [] + failures: list[dict[str, str]] = [] + attempt_count = 0 + fatal_error: str | None = None + readiness_written = False + + # A host CPU percentage requires a delta between two /proc/stat snapshots. + next_sample_at = started_monotonic + interval_seconds try: while RUNNING: + if host_collector is not None: + wait_until(next_sample_at) + if not RUNNING: + break + attempt_count += 1 try: - sample = collect_once(container_name) - SAMPLES.append(sample) + if host_collector is not None: + samples.append(host_collector.collect()) + elif injected_single_sample: + samples.append(collect_container_once(target)) + else: + assert container_stream is not None + samples.append(container_stream.collect()) + except CollectionStopped: + attempt_count -= 1 + break except Exception as exc: - print(f"Warning: failed to collect stats: {exc}", file=sys.stderr) - - slept = 0.0 - while RUNNING and slept < interval_seconds: - chunk = min(0.2, interval_seconds - slept) - time.sleep(chunk) - slept += chunk + failure = {"timestamp": utc_now(), "error": str(exc)} + failures.append(failure) + print( + f"Resource collection attempt {attempt_count} failed: {exc}", + file=sys.stderr, + ) + # A failed sample already makes the run ineligible, so continuing + # can only create duplicate errors and cannot salvage the audit. + break + + if samples and not readiness_written: + try: + write_json( + readiness_path_for(output_path), + { + "schema_version": 1, + "status": "ready", + "target": target, + "target_type": target_type, + "ready_at": utc_now(), + "baseline_timestamp": samples[0]["timestamp"], + }, + ) + readiness_written = True + except Exception as exc: + fatal_error = f"failed to write collector readiness marker: {exc}" + print(fatal_error, file=sys.stderr) + break + + if host_collector is not None: + next_sample_at += interval_seconds + if next_sample_at <= time.monotonic(): + # Do not issue a burst of catch-up samples. The validator will + # reject the resulting sparse coverage instead of hiding it. + next_sample_at = time.monotonic() + interval_seconds except Exception as exc: + fatal_error = str(exc) print(f"Fatal error in collector loop: {exc}", file=sys.stderr) - exit_code = 1 finally: - try: - write_samples() - except Exception as exc: - print(f"Failed to write output JSON: {exc}", file=sys.stderr) - return 1 + if container_stream is not None: + container_stream.close() + + finished_monotonic = time.monotonic() + finished_at = utc_now() + succeeded = ( + bool(samples) + and readiness_written + and not failures + and fatal_error is None + and STOP_SIGNAL is not None + ) + audit: dict[str, Any] = { + "schema_version": 1, + "status": "complete" if succeeded else "failed", + "target": target, + "target_type": target_type, + "output_path": str(output_path), + "audit_path": str(audit_path_for(output_path)), + "readiness_path": str(readiness_path_for(output_path)), + "interval_seconds": interval_seconds, + "started_at": started_at, + "finished_at": finished_at, + "elapsed_seconds": finished_monotonic - started_monotonic, + "termination": { + "reason": "signal" if STOP_SIGNAL is not None else "loop_exit", + "signal": STOP_SIGNAL, + }, + "attempt_count": attempt_count, + "sample_count": len(samples), + "failure_count": len(failures), + "failures": failures, + "fatal_error": fatal_error, + } + if host_collector is not None: + audit["host"] = host_collector.audit_metadata() + else: + audit["source"] = "persistent docker stats stream" + + write_failed = False + try: + # Keep the historical top-level list schema for downstream consumers. + write_json(output_path, samples) + except Exception as exc: + write_failed = True + print(f"Failed to write sample JSON: {exc}", file=sys.stderr) + try: + write_json(audit_path_for(output_path), audit) + except Exception as exc: + write_failed = True + print(f"Failed to write audit JSON: {exc}", file=sys.stderr) - return exit_code + return 0 if succeeded and not write_failed else 1 if __name__ == "__main__": diff --git a/benchmark/cpp/CMakeLists.txt b/benchmark/cpp/CMakeLists.txt index 09c71f8..f860086 100644 --- a/benchmark/cpp/CMakeLists.txt +++ b/benchmark/cpp/CMakeLists.txt @@ -28,7 +28,7 @@ target_link_libraries(benchmark-server PRIVATE mcp-cpp-sdk) if(CMAKE_BUILD_TYPE MATCHES Release) if(CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang") - target_compile_options(benchmark-server PRIVATE -O3 -march=native) + target_compile_options(benchmark-server PRIVATE -O3) elseif(MSVC) target_compile_options(benchmark-server PRIVATE /O2) endif() diff --git a/benchmark/cpp/benchmark_server.cpp b/benchmark/cpp/benchmark_server.cpp index 5b5721c..d621b7e 100644 --- a/benchmark/cpp/benchmark_server.cpp +++ b/benchmark/cpp/benchmark_server.cpp @@ -806,8 +806,6 @@ int main() { StreamableHttpSessionManager manager(io_ctx.get_executor(), "0.0.0.0", mcp_port, create_server_with_tools); - manager.set_stateless_json_mode(true); - manager.set_custom_request_handler( [](const boost::beast::http::request& req) -> std::optional> { diff --git a/benchmark/docker-compose.yml b/benchmark/docker-compose.yml index d71c305..fb61d56 100644 --- a/benchmark/docker-compose.yml +++ b/benchmark/docker-compose.yml @@ -1,7 +1,9 @@ services: redis: - image: redis:7-alpine + image: redis@sha256:8b81dd37ff027bec4e516d41acfbe9fe2460070dc6d4a4570a2ac5b9d59df065 container_name: mcp-redis + cpuset: "${BENCHMARK_CPUSET:-}" + command: redis-server --save "" --appendonly no ports: - "6379:6379" deploy: @@ -20,6 +22,7 @@ services: api-service: build: ./benchmark-mcp-servers-v2/api-service container_name: mcp-api-service + cpuset: "${BENCHMARK_CPUSET:-}" ports: - "8100:8100" deploy: @@ -36,7 +39,10 @@ services: - mcp-network redis-seeder: - build: ./benchmark-mcp-servers-v2/infra/redis + build: + context: . + dockerfile: redis-seeder.Dockerfile + image: mcp-benchmark-redis-seeder:python-redis-5.2.1 container_name: mcp-redis-seeder profiles: ["seeder"] depends_on: @@ -45,6 +51,9 @@ services: environment: - REDIS_HOST=redis - REDIS_PORT=6379 + volumes: + - ./benchmark-mcp-servers-v2/infra/redis/seed.py:/seed.py:ro + command: ["/seed.py"] networks: - mcp-network @@ -53,6 +62,7 @@ services: context: .. dockerfile: benchmark/cpp/Dockerfile container_name: mcp-cpp-server + cpuset: "${BENCHMARK_CPUSET:-}" ports: - "8080:8080" environment: @@ -76,14 +86,40 @@ services: networks: - mcp-network + ours-comparable-server: + build: + context: .. + dockerfile: benchmark/alternatives/Dockerfile.ours + container_name: mcp-ours-comparable-server + cpuset: "${BENCHMARK_CPUSET:-}" + ports: + - "8089:8080" + environment: + - PORT=8080 + - REDIS_URL=redis://redis:6379 + - API_SERVICE_URL=http://api-service:8100 + deploy: + resources: + limits: + cpus: "2.0" + memory: 2G + depends_on: + redis: + condition: service_healthy + api-service: + condition: service_healthy + networks: + - mcp-network + # Upstream servers for comparison python-server: build: context: ./benchmark-mcp-servers-v2/python-server dockerfile: Dockerfile container_name: mcp-python-server + cpuset: "${BENCHMARK_CPUSET:-}" ports: - - "8081:8081" + - "8081:8082" environment: - REDIS_URL=redis://redis:6379 - API_SERVICE_URL=http://api-service:8100 @@ -98,7 +134,7 @@ services: api-service: condition: service_healthy healthcheck: - test: ["CMD", "python", "-c", "import urllib.request,json; urllib.request.urlopen(urllib.request.Request('http://localhost:8081/mcp',data=json.dumps({'jsonrpc':'2.0','id':1,'method':'initialize','params':{'protocolVersion':'2024-11-05','capabilities':{},'clientInfo':{'name':'health','version':'1.0'}}}).encode(),headers={'Content-Type':'application/json','Accept':'application/json, text/event-stream'}),timeout=5)"] + test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://localhost:8082/health', timeout=5)"] interval: 10s timeout: 5s retries: 3 @@ -110,8 +146,9 @@ services: context: ./benchmark-mcp-servers-v2/go-server dockerfile: Dockerfile container_name: mcp-go-server + cpuset: "${BENCHMARK_CPUSET:-}" ports: - - "8082:8082" + - "8082:8081" environment: - REDIS_URL=redis://redis:6379 - API_SERVICE_URL=http://api-service:8100 @@ -126,7 +163,7 @@ services: api-service: condition: service_healthy healthcheck: - test: ["CMD", "wget", "--no-verbose", "--tries=1", "--spider", "http://localhost:8082/health"] + test: ["CMD", "wget", "--no-verbose", "--tries=1", "--spider", "http://localhost:8081/health"] interval: 10s timeout: 5s retries: 3 @@ -138,8 +175,9 @@ services: context: ./benchmark-mcp-servers-v2/rust-server dockerfile: Dockerfile container_name: mcp-rust-server + cpuset: "${BENCHMARK_CPUSET:-}" ports: - - "8083:8083" + - "8083:8095" environment: - REDIS_URL=redis://redis:6379 - API_SERVICE_URL=http://api-service:8100 @@ -154,7 +192,7 @@ services: api-service: condition: service_healthy healthcheck: - test: ["CMD", "curl", "-sf", "http://localhost:8083/health"] + test: ["CMD", "curl", "-sf", "http://localhost:8095/health"] interval: 10s timeout: 5s retries: 5 @@ -162,6 +200,113 @@ services: networks: - mcp-network + hkr04-server: + build: + context: . + dockerfile: alternatives/Dockerfile + args: + SDK_NAME: hkr04 + container_name: mcp-hkr04-server + cpuset: "${BENCHMARK_CPUSET:-}" + ports: + - "8084:8080" + environment: + - PORT=8080 + - REDIS_URL=redis://redis:6379 + - API_SERVICE_URL=http://api-service:8100 + deploy: + resources: + limits: + cpus: "2.0" + memory: 2G + depends_on: + redis: + condition: service_healthy + api-service: + condition: service_healthy + networks: + - mcp-network + + fastmcpp-server: + build: + context: . + dockerfile: alternatives/Dockerfile + args: + SDK_NAME: fastmcpp + container_name: mcp-fastmcpp-server + cpuset: "${BENCHMARK_CPUSET:-}" + ports: + - "8085:8080" + environment: + - PORT=8080 + - REDIS_URL=redis://redis:6379 + - API_SERVICE_URL=http://api-service:8100 + deploy: + resources: + limits: + cpus: "2.0" + memory: 2G + depends_on: + redis: + condition: service_healthy + api-service: + condition: service_healthy + networks: + - mcp-network + + cxxmcp-server: + build: + context: . + dockerfile: alternatives/Dockerfile + args: + SDK_NAME: cxxmcp + container_name: mcp-cxxmcp-server + cpuset: "${BENCHMARK_CPUSET:-}" + ports: + - "8086:8080" + environment: + - PORT=8080 + - REDIS_URL=redis://redis:6379 + - API_SERVICE_URL=http://api-service:8100 + deploy: + resources: + limits: + cpus: "2.0" + memory: 2G + depends_on: + redis: + condition: service_healthy + api-service: + condition: service_healthy + networks: + - mcp-network + + neumann-server: + build: + context: . + dockerfile: alternatives/Dockerfile + args: + SDK_NAME: neumann + container_name: mcp-neumann-server + cpuset: "${BENCHMARK_CPUSET:-}" + ports: + - "8087:8080" + environment: + - PORT=8080 + - REDIS_URL=redis://redis:6379 + - API_SERVICE_URL=http://api-service:8100 + deploy: + resources: + limits: + cpus: "2.0" + memory: 2G + depends_on: + redis: + condition: service_healthy + api-service: + condition: service_healthy + networks: + - mcp-network networks: mcp-network: diff --git a/benchmark/redis-seeder.Dockerfile b/benchmark/redis-seeder.Dockerfile new file mode 100644 index 0000000..003b9e2 --- /dev/null +++ b/benchmark/redis-seeder.Dockerfile @@ -0,0 +1,5 @@ +FROM python@sha256:57cd7c3a7a273101a6485ba99423ee568157882804b1124b4dd04266317710de + +RUN pip install --disable-pip-version-check --no-cache-dir redis==5.2.1 + +ENTRYPOINT ["python3"] diff --git a/benchmark/run.sh b/benchmark/run.sh index fe61095..2443da0 100755 --- a/benchmark/run.sh +++ b/benchmark/run.sh @@ -3,140 +3,439 @@ set -euo pipefail SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" PROJECT_DIR="$(dirname "$SCRIPT_DIR")" -RESULTS_DIR="$SCRIPT_DIR/results/$(date +%Y%m%d_%H%M%S)" +COMPOSE_FILE="$SCRIPT_DIR/docker-compose.yml" +RUNTIME_COMPOSE_FILE="$COMPOSE_FILE" UPSTREAM_BENCHMARK_DIR="$SCRIPT_DIR/benchmark-mcp-servers-v2" UPSTREAM_BENCHMARK_REPO="https://github.com/thiagomendes/benchmark-mcp-servers-v2.git" UPSTREAM_BENCHMARK_COMMIT="8a9a5f8ef505f46b6079072ef4603304ca672e33" -UPSTREAM_K6_SCRIPT="$UPSTREAM_BENCHMARK_DIR/benchmark/benchmark.js" +K6_SCRIPT="$SCRIPT_DIR/alternatives/benchmark.js" +K6_RUNTIME_SCRIPT_DIR="$SCRIPT_DIR/alternatives" +RUNTIME_COMPOSE_SHA256="" +RUNTIME_K6_SHA256="" +ALTERNATIVE_MANIFEST="$SCRIPT_DIR/alternatives/sources.tsv" +ALTERNATIVE_SOURCE_DIR="$SCRIPT_DIR/alternative-sdks" +MCP_PROTOCOL_VERSION="2024-11-05" +ELIGIBILITY_CONTRACT="upstream-v2-strict-mcp-v1" +SUPPLEMENTAL_CONTRACT="adapter-exact-v1" +HOST_NETWORK_POLICY="baseline-cohort-retire-v1" +K6_IMAGE="grafana/k6@sha256:82e44a45a38ed22bf5636fe50fe8a07967c3074f7aa66567c6a7501ab9bb3a9f" + +BENCHMARK_RUNS="${BENCHMARK_RUNS:-3}" +BENCHMARK_VUS="${BENCHMARK_VUS:-50}" +BENCHMARK_RAMP_DURATION="${BENCHMARK_RAMP_DURATION:-15s}" +BENCHMARK_WARMUP_DURATION="${BENCHMARK_WARMUP_DURATION:-60s}" +BENCHMARK_MEASURE_DURATION="${BENCHMARK_MEASURE_DURATION:-5m}" +BENCHMARK_ORDER_SEED="${BENCHMARK_ORDER_SEED:-$(date +%s)}" +# The load-generator contract is fixed so a run cannot be labeled publishable +# while silently using limits that differ from the recorded expectations. +readonly BENCHMARK_K6_CPUS=4 +readonly BENCHMARK_K6_MEMORY="2g" +readonly BENCHMARK_K6_MEMORY_BYTES=2147483648 +readonly BENCHMARK_REQUEST_CONCURRENCY=50 +readonly BENCHMARK_LOCK_FILE="/tmp/mcp-cpp-sdk-benchmark.lock" +RESULTS_DIR="" +RUN_SUCCEEDED=0 +BENCHMARK_ACTIVE=0 +CURRENT_SERVER="" +CURRENT_SERVER_RESULTS="" +CURRENT_K6_CONTAINER="" +HOST_CPU_LIMIT="" +HOST_MEMORY_BYTES="" +HOST_CPUSET="" +FAILURE_STAGE="" +FAILURE_REASON="" +FAILURE_EVIDENCE_LOG="" RED='\033[0;31m' GREEN='\033[0;32m' -YELLOW='\033[1;33m' BLUE='\033[0;34m' NC='\033[0m' -info() { - printf "%b[INFO]%b %s\n" "$BLUE" "$NC" "$*" -} - -ok() { - printf "%b[OK]%b %s\n" "$GREEN" "$NC" "$*" -} +info() { printf "%b[INFO]%b %s\n" "$BLUE" "$NC" "$*"; } +ok() { printf "%b[OK]%b %s\n" "$GREEN" "$NC" "$*"; } +error() { printf "%b[ERROR]%b %s\n" "$RED" "$NC" "$*" >&2; } -warn() { - printf "%b[WARN]%b %s\n" "$YELLOW" "$NC" "$*" -} +acquire_benchmark_lock() { + command -v flock >/dev/null || { + error "flock not found; cannot protect the fixed benchmark ports and containers" + return 1 + } -error() { - printf "%b[ERROR]%b %s\n" "$RED" "$NC" "$*" >&2 + exec 9>> "$BENCHMARK_LOCK_FILE" + if ! flock -n 9; then + local holder="" + holder="$(head -n 1 "$BENCHMARK_LOCK_FILE" 2>/dev/null || true)" + error "Another benchmark invocation is active (PID ${holder:-unknown})" + return 1 + fi + printf '%s\n' "$$" > "$BENCHMARK_LOCK_FILE" } declare -A SERVICE_NAME=( [cpp]="cpp-server" + [ours-comparable]="ours-comparable-server" [python]="python-server" [go]="go-server" [rust]="rust-server" + [hkr04]="hkr04-server" + [fastmcpp]="fastmcpp-server" + [cxxmcp]="cxxmcp-server" + [neumann]="neumann-server" ) declare -A CONTAINER_NAME=( [cpp]="mcp-cpp-server" + [ours-comparable]="mcp-ours-comparable-server" [python]="mcp-python-server" [go]="mcp-go-server" [rust]="mcp-rust-server" + [hkr04]="mcp-hkr04-server" + [fastmcpp]="mcp-fastmcpp-server" + [cxxmcp]="mcp-cxxmcp-server" + [neumann]="mcp-neumann-server" ) declare -A MCP_URL=( [cpp]="http://localhost:8080/mcp" + [ours-comparable]="http://localhost:8089/mcp" [python]="http://localhost:8081/mcp" [go]="http://localhost:8082/mcp" [rust]="http://localhost:8083/mcp" + [hkr04]="http://localhost:8084/mcp" + [fastmcpp]="http://localhost:8085/mcp" + [cxxmcp]="http://localhost:8086/mcp" + [neumann]="http://localhost:8087/mcp" ) declare -A HEALTH_URL=( [cpp]="http://localhost:8080/health" + [ours-comparable]="http://localhost:8089/mcp" [python]="http://localhost:8081/health" [go]="http://localhost:8082/health" [rust]="http://localhost:8083/health" + [hkr04]="http://localhost:8084/mcp" + [fastmcpp]="http://localhost:8085/mcp" + [cxxmcp]="http://localhost:8086/mcp" + [neumann]="http://localhost:8087/mcp" ) -stats_pid="" +declare -A HEALTH_KIND=( + [cpp]="get" + [ours-comparable]="mcp" + [python]="get" + [go]="get" + [rust]="get" + [hkr04]="mcp" + [fastmcpp]="mcp" + [cxxmcp]="mcp" + [neumann]="mcp" +) -cleanup_on_exit() { - if [[ -n "${stats_pid:-}" ]] && kill -0 "$stats_pid" 2>/dev/null; then - warn "Stopping running stats collector (PID $stats_pid)..." - kill "$stats_pid" 2>/dev/null || true - wait "$stats_pid" 2>/dev/null || true - fi -} +declare -A EXPECTED_SERVER_TYPE=( + [cpp]="cpp" + [ours-comparable]="cpp-sdk" + [hkr04]="cpp-sdk" + [fastmcpp]="cpp-sdk" + [cxxmcp]="cpp-sdk" + [neumann]="cpp-sdk" + [python]="python" + [go]="go" + [rust]="rust" +) -trap cleanup_on_exit EXIT +ALL_MCP_SERVICES=( + cpp-server ours-comparable-server python-server go-server rust-server + hkr04-server fastmcpp-server cxxmcp-server neumann-server +) +CPP_SDK_SERVERS=(ours-comparable hkr04 fastmcpp cxxmcp neumann) +BASELINE_SERVERS=(python go rust) + +declare -A REQUIRE_SUPPLEMENTAL=( + [ours-comparable]=1 + [hkr04]=1 + [fastmcpp]=1 + [cxxmcp]=1 + [neumann]=1 +) + +declare -A ALTERNATIVE_REPO=() +declare -A ALTERNATIVE_COMMIT=() +declare -A ALTERNATIVE_STATUS=() +collector_pids=() +collector_ready_files=() usage() { cat <<'USAGE' Usage: - ./run.sh # benchmark cpp, python, go, rust - ./run.sh all # benchmark cpp, python, go, rust - ./run.sh cpp # benchmark only cpp - ./run.sh cpp python # benchmark selected servers - ./run.sh cpp,python # benchmark selected servers (comma-separated) + ./run.sh # five comparable C++ SDKs + ./run.sh cpp-sdks # five comparable C++ SDKs + ./run.sh baseline # Python, Go, and Rust baselines + ./run.sh all # all corrected C++ and language baselines + ./run.sh python go rust # selected servers + ./run.sh hkr04,cxxmcp # selected servers (comma-separated) + +Production profile: 3 rounds, 50 VUs, 15s ramp + 60s warmup, then a +separate constant-50-VU 5m measurement. Any override is labeled smoke and +is never a publishable result. USAGE } -ensure_upstream_benchmark_servers() { - command -v git >/dev/null || { error "git not found (required to clone upstream benchmark servers)"; return 1; } +verify_runtime_harness() { + [[ -n "$RUNTIME_COMPOSE_SHA256" && -n "$RUNTIME_K6_SHA256" ]] || return 0 + local compose_sha256 k6_sha256 + compose_sha256="$(sha256sum "$RUNTIME_COMPOSE_FILE" | awk '{print $1}')" + k6_sha256="$(sha256sum "$K6_RUNTIME_SCRIPT_DIR/benchmark.js" | awk '{print $1}')" + [[ "$compose_sha256" == "$RUNTIME_COMPOSE_SHA256" \ + && "$k6_sha256" == "$RUNTIME_K6_SHA256" ]] || { + error "Immutable runtime harness changed during the benchmark" + return 1 + } +} - if [[ ! -d "$UPSTREAM_BENCHMARK_DIR/.git" ]]; then - if [[ -e "$UPSTREAM_BENCHMARK_DIR" ]]; then - error "Expected clone target exists but is not a git repo: $UPSTREAM_BENCHMARK_DIR" - error "Please remove it or convert it into a valid clone of $UPSTREAM_BENCHMARK_REPO" - return 1 +compose() { + verify_runtime_harness + docker compose --project-directory "$SCRIPT_DIR" -f "$RUNTIME_COMPOSE_FILE" "$@" +} + +stop_collectors() { + local pid + local failed=0 + for pid in "${collector_pids[@]:-}"; do + if [[ -n "$pid" ]] && kill -0 "$pid" 2>/dev/null; then + kill "$pid" 2>/dev/null || true + fi + done + for pid in "${collector_pids[@]:-}"; do + if [[ -n "$pid" ]]; then + wait "$pid" 2>/dev/null || failed=1 fi + done + collector_pids=() + collector_ready_files=() + return "$failed" +} - info "Cloning upstream benchmark servers to $UPSTREAM_BENCHMARK_DIR" - git clone "$UPSTREAM_BENCHMARK_REPO" "$UPSTREAM_BENCHMARK_DIR" - ok "Upstream benchmark servers cloned" +finalize_manifest() { + local exit_code="$1" + [[ -n "$RESULTS_DIR" && -f "$RESULTS_DIR/run_manifest.json" ]] || return 0 + local status="failed" + if [[ "$exit_code" -eq 0 && "$RUN_SUCCEEDED" -eq 1 ]]; then + status="complete" fi + local temporary="$RESULTS_DIR/run_manifest.json.tmp" + jq --arg status "$status" --arg completed_at "$(date -u +%Y-%m-%dT%H:%M:%SZ)" \ + --arg failure_stage "$FAILURE_STAGE" \ + --arg failure_reason "$FAILURE_REASON" \ + --arg failure_log "$FAILURE_EVIDENCE_LOG" \ + --argjson exit_code "$exit_code" \ + '.status = $status | .completed_at = $completed_at + | .publishable_candidate = ($status == "complete" and .profile == "production") + | if $status == "failed" then + .failure = { + exit_code: $exit_code, + stage: (if $failure_stage == "" then "unknown" else $failure_stage end), + reason: (if $failure_reason == "" then null else $failure_reason end), + evidence_log: (if $failure_log == "" then null else $failure_log end) + } + else del(.failure) + end' \ + "$RESULTS_DIR/run_manifest.json" > "$temporary" \ + && mv "$temporary" "$RESULTS_DIR/run_manifest.json" +} - if [[ -n "$(git -C "$UPSTREAM_BENCHMARK_DIR" status --porcelain --untracked-files=no)" ]]; then - error "Upstream benchmark repo has local modifications: $UPSTREAM_BENCHMARK_DIR" - error "Please clean it before running benchmark to keep pinned reproducibility" - return 1 +clear_failure_context() { + FAILURE_STAGE="" + FAILURE_REASON="" + FAILURE_EVIDENCE_LOG="" +} + +record_failure_from_log() { + local stage="$1" + local log_file="$2" + local fallback_reason="$3" + local final_line="" + + FAILURE_STAGE="$stage" + FAILURE_EVIDENCE_LOG="${log_file#"$RESULTS_DIR"/}" + if [[ -f "$log_file" ]]; then + final_line="$(awk 'NF { line = $0 } END { print line }' "$log_file")" fi + FAILURE_REASON="${final_line:-$fallback_reason}" +} - info "Pinning upstream benchmark repo to commit $UPSTREAM_BENCHMARK_COMMIT" - git -C "$UPSTREAM_BENCHMARK_DIR" fetch --depth 1 origin "$UPSTREAM_BENCHMARK_COMMIT" - git -C "$UPSTREAM_BENCHMARK_DIR" checkout --detach "$UPSTREAM_BENCHMARK_COMMIT" +record_collector_failure() { + local stage="$1" + local log_file="$2" + local fallback_reason="$3" + shift 3 + local structured_reason="" + + FAILURE_STAGE="$stage" + FAILURE_EVIDENCE_LOG="${log_file#"$RESULTS_DIR"/}" + if structured_reason="$( + python3 "$SCRIPT_DIR/summarize_collector_failures.py" "$@" 2>/dev/null + )" && [[ -n "$structured_reason" ]]; then + FAILURE_REASON="$structured_reason" + return + fi + record_failure_from_log "$stage" "$log_file" "$fallback_reason" +} - local current_commit - current_commit="$(git -C "$UPSTREAM_BENCHMARK_DIR" rev-parse HEAD)" - if [[ "$current_commit" != "$UPSTREAM_BENCHMARK_COMMIT" ]]; then - error "Failed to pin upstream benchmark repo to expected commit" - error "Current: $current_commit" - error "Expected: $UPSTREAM_BENCHMARK_COMMIT" - return 1 +record_resource_collection_failure() { + local stage="$1" + local server_results="$2" + local run_idx="$3" + local fallback_reason="$4" + local collector_log="$server_results/resource_collection_run${run_idx}.log" + local collector_file + + { + for collector_file in \ + "$server_results/stats_run${run_idx}.log" \ + "$server_results/redis_stats_run${run_idx}.log" \ + "$server_results/api_stats_run${run_idx}.log" \ + "$server_results/k6_stats_run${run_idx}.log" \ + "$server_results/host_stats_run${run_idx}.log"; do + if [[ -s "$collector_file" ]]; then + printf '[%s]\n' "$(basename "$collector_file")" + tail -n 20 "$collector_file" + fi + done + } > "$collector_log" + + record_collector_failure "$stage" "$collector_log" "$fallback_reason" \ + "$server_results/stats_run${run_idx}.audit.json" \ + "$server_results/redis_stats_run${run_idx}.audit.json" \ + "$server_results/api_stats_run${run_idx}.audit.json" \ + "$server_results/k6_stats_run${run_idx}.audit.json" \ + "$server_results/host_stats_run${run_idx}.audit.json" +} + +run_protocol_verifier() { + local stage="$1" + local log_file="$2" + shift 2 + local -a pipeline_status + local verifier_exit tee_exit + + set +e + python3 "$SCRIPT_DIR/verify_server.py" "$@" 2>&1 | tee "$log_file" + pipeline_status=("${PIPESTATUS[@]}") + set -e + verifier_exit="${pipeline_status[0]}" + tee_exit="${pipeline_status[1]}" + if [[ "$verifier_exit" -eq 0 && "$tee_exit" -eq 0 ]]; then + return 0 fi - local required_paths=( - "$UPSTREAM_BENCHMARK_DIR/api-service" - "$UPSTREAM_BENCHMARK_DIR/infra/redis" - "$UPSTREAM_BENCHMARK_DIR/benchmark" - "$UPSTREAM_BENCHMARK_DIR/python-server" - "$UPSTREAM_BENCHMARK_DIR/go-server" - "$UPSTREAM_BENCHMARK_DIR/rust-server" - ) - for p in "${required_paths[@]}"; do - if [[ ! -d "$p" ]]; then - error "Missing required upstream path after checkout: $p" - return 1 + FAILURE_STAGE="$stage" + FAILURE_EVIDENCE_LOG="${log_file#"$RESULTS_DIR"/}" + FAILURE_REASON="$(tail -n 1 "$log_file" 2>/dev/null || true)" + if [[ "$verifier_exit" -ne 0 ]]; then + return "$verifier_exit" + fi + return "$tee_exit" +} + +cleanup_on_exit() { + local exit_code=$? + trap - EXIT + if [[ "$BENCHMARK_ACTIVE" -eq 1 ]]; then + stop_collectors || true + if [[ -n "$CURRENT_K6_CONTAINER" ]]; then + docker rm -f "$CURRENT_K6_CONTAINER" >/dev/null 2>&1 || true fi - done + if [[ -n "$CURRENT_SERVER" && -n "$CURRENT_SERVER_RESULTS" ]]; then + docker logs "$CURRENT_SERVER" > "$CURRENT_SERVER_RESULTS/server_failure.log" 2>&1 || true + fi + compose stop "${ALL_MCP_SERVICES[@]}" >/dev/null 2>&1 || true + compose stop redis api-service >/dev/null 2>&1 || true + fi + if ! finalize_manifest "$exit_code"; then + error "Failed to finalize benchmark run manifest" + if [[ "$exit_code" -eq 0 ]]; then + exit_code=1 + fi + fi + exit "$exit_code" +} + +trap cleanup_on_exit EXIT - if [[ ! -f "$UPSTREAM_K6_SCRIPT" ]]; then - error "Missing required upstream k6 script: $UPSTREAM_K6_SCRIPT" +load_alternative_manifest() { + [[ -f "$ALTERNATIVE_MANIFEST" ]] || { + error "Missing alternative source manifest: $ALTERNATIVE_MANIFEST" return 1 + } + while IFS=$'\t' read -r id repository commit status; do + [[ -z "$id" || "$id" == \#* ]] && continue + ALTERNATIVE_REPO["$id"]="$repository" + ALTERNATIVE_COMMIT["$id"]="$commit" + ALTERNATIVE_STATUS["$id"]="$status" + done < "$ALTERNATIVE_MANIFEST" +} + +ensure_git_checkout() { + local name="$1" + local repository="$2" + local commit="$3" + local target="$4" + local fresh_clone=0 + + if [[ ! -d "$target/.git" ]]; then + [[ ! -e "$target" ]] || { + error "Clone target exists but is not a git repository: $target" + return 1 + } + info "Cloning $name from $repository" + git clone --filter=blob:none --no-checkout "$repository" "$target" + fresh_clone=1 + fi + local actual_origin + actual_origin="$(git -C "$target" remote get-url origin)" + [[ "$actual_origin" == "$repository" ]] || { + error "$name origin mismatch: expected $repository, found $actual_origin" + return 1 + } + if [[ "$fresh_clone" -eq 0 && -n "$(git -C "$target" status --porcelain)" ]]; then + error "$name source has local modifications: $target" + return 1 + fi + if [[ "$fresh_clone" -eq 1 || "$(git -C "$target" rev-parse HEAD)" != "$commit" ]]; then + git -C "$target" fetch --depth 1 origin "$commit" + git -C "$target" checkout --detach "$commit" fi + [[ "$(git -C "$target" rev-parse HEAD)" == "$commit" ]] || { + error "Failed to pin $name to $commit" + return 1 + } + [[ -z "$(git -C "$target" status --porcelain)" ]] || { + error "$name source is not clean after pinning" + return 1 + } + ok "Using $name pinned at $commit" +} + +ensure_upstream_sources() { + ensure_git_checkout \ + "upstream benchmark" "$UPSTREAM_BENCHMARK_REPO" \ + "$UPSTREAM_BENCHMARK_COMMIT" "$UPSTREAM_BENCHMARK_DIR" + local required_paths=(api-service infra/redis python-server go-server rust-server) + local relative + for relative in "${required_paths[@]}"; do + [[ -d "$UPSTREAM_BENCHMARK_DIR/$relative" ]] || { + error "Missing upstream path: $UPSTREAM_BENCHMARK_DIR/$relative" + return 1 + } + done +} - ok "Using upstream benchmark servers repo pinned at $UPSTREAM_BENCHMARK_COMMIT" +ensure_alternative_source() { + local id="$1" + mkdir -p "$ALTERNATIVE_SOURCE_DIR" + ensure_git_checkout \ + "$id" "${ALTERNATIVE_REPO[$id]}" "${ALTERNATIVE_COMMIT[$id]}" \ + "$ALTERNATIVE_SOURCE_DIR/$id" } wait_for_http() { @@ -145,11 +444,28 @@ wait_for_http() { local timeout_secs="$3" local start_ts start_ts="$(date +%s)" + info "Waiting for $name at $url" + until curl -fsS "$url" >/dev/null 2>&1; do + if (( $(date +%s) - start_ts >= timeout_secs )); then + error "Timeout waiting for $name at $url" + return 1 + fi + sleep 1 + done + ok "$name is healthy" +} - info "Waiting for $name at $url (timeout ${timeout_secs}s)..." +wait_for_mcp() { + local name="$1" + local url="$2" + local timeout_secs="$3" + local start_ts + start_ts="$(date +%s)" + info "Waiting for $name MCP initialize at $url" while true; do - if curl -fsS "$url" >/dev/null 2>&1; then - ok "$name is healthy" + if python3 "$SCRIPT_DIR/verify_server.py" "$url" \ + --name "$name-readiness" --initialize-only >/dev/null 2>&1; then + ok "$name is accepting MCP requests" return 0 fi if (( $(date +%s) - start_ts >= timeout_secs )); then @@ -160,248 +476,672 @@ wait_for_http() { done } -warmup_single_session() { - local server_url="$1" - local tool_name="$2" - local args_json="$3" - - local hdr - hdr="$(mktemp)" - - curl -fsS -D "$hdr" -o /dev/null -X POST "$server_url" \ - -H "Content-Type: application/json" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ - -d '{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"benchmark-warmup","version":"1.0"}}}' - - local session_id - session_id="$(tr -d '\r' < "$hdr" | grep -i '^mcp-session-id:' | sed 's/^[^:]*: *//' | head -n1 || true)" - rm -f "$hdr" - - curl -fsS -o /dev/null -X POST "$server_url" \ - -H "Content-Type: application/json" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ - ${session_id:+-H "Mcp-Session-Id: $session_id"} \ - -d '{"jsonrpc":"2.0","method":"notifications/initialized"}' - - curl -fsS -o /dev/null -X POST "$server_url" \ - -H "Content-Type: application/json" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ - ${session_id:+-H "Mcp-Session-Id: $session_id"} \ - -d "{\"jsonrpc\":\"2.0\",\"id\":2,\"method\":\"tools/call\",\"params\":{\"name\":\"${tool_name}\",\"arguments\":${args_json}}}" - - if [[ -n "$session_id" ]]; then - curl -fsS -o /dev/null -X DELETE "$server_url" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ - -H "Mcp-Session-Id: $session_id" - fi +reset_redis_dataset() { + compose exec -T redis redis-cli FLUSHDB >/dev/null + docker rm -f mcp-redis-seeder >/dev/null 2>&1 || true + compose --profile seeder run --rm redis-seeder >/dev/null } -run_warmup() { - local server_name="$1" - local server_url="$2" +start_collector() { + local container="$1" + local output="$2" + python3 "$SCRIPT_DIR/collect_stats.py" "$container" "$output" 1.0 \ + > "${output%.json}.log" 2>&1 & + collector_pids+=("$!") + collector_ready_files+=("${output%.json}.ready.json") +} + +wait_for_collectors_ready() { + local timeout_secs="${1:-15}" + local started_at + local index + local all_ready + started_at="$(date +%s)" + while true; do + all_ready=1 + for index in "${!collector_pids[@]}"; do + if ! kill -0 "${collector_pids[$index]}" 2>/dev/null; then + error "Resource collector exited before readiness: ${collector_ready_files[$index]}" + return 1 + fi + if ! jq -e '.schema_version == 1 and .status == "ready"' \ + "${collector_ready_files[$index]}" >/dev/null 2>&1; then + all_ready=0 + fi + done + if [[ "$all_ready" -eq 1 ]]; then + # Close the small ready-marker/process-exit race before measurement. + for index in "${!collector_pids[@]}"; do + kill -0 "${collector_pids[$index]}" 2>/dev/null || { + error "Resource collector exited at readiness: ${collector_ready_files[$index]}" + return 1 + } + done + return 0 + fi + if (( $(date +%s) - started_at >= timeout_secs )); then + error "Timed out waiting for resource collectors to capture baselines" + return 1 + fi + sleep 0.1 + done +} - info "Warmup for $server_name: 5 initialize requests" - for _ in {1..5}; do - local hdr - hdr="$(mktemp)" +wait_for_container() { + local container="$1" + local attempt + for attempt in $(seq 1 100); do + if docker inspect "$container" >/dev/null 2>&1; then + return 0 + fi + sleep 0.1 + done + error "Container did not start: $container" + return 1 +} - curl -fsS -D "$hdr" -o /dev/null -X POST "$server_url" \ +resume_k6() { + local server="$1" + local run_idx="$2" + local attempt + local payload='{"data":{"type":"status","id":"default","attributes":{"paused":false}}}' + for attempt in $(seq 1 100); do + if curl -fsS -X PATCH "http://127.0.0.1:6565/v1/status" \ -H "Content-Type: application/json" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ - -d '{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-11-25","capabilities":{},"clientInfo":{"name":"benchmark-warmup","version":"1.0"}}}' - - local session_id - session_id="$(tr -d '\r' < "$hdr" | grep -i '^mcp-session-id:' | sed 's/^[^:]*: *//' | head -n1 || true)" - rm -f "$hdr" - - if [[ -n "$session_id" ]]; then - curl -fsS -o /dev/null -X DELETE "$server_url" \ - -H "Accept: application/json, text/event-stream" \ - -H "MCP-Protocol-Version: 2025-11-25" \ - -H "Mcp-Session-Id: $session_id" + --data "$payload" >/dev/null 2>&1; then + return 0 fi + sleep 0.1 done + error "[$server] could not resume paused k6 run $run_idx" + return 1 +} + +run_k6_warmup() { + local server="$1" + local server_url="$2" + local server_results="$3" + local run_idx="$4" + local negotiated_protocol_version="$5" + verify_runtime_harness + local container="mcp-k6-warmup-${server//[^a-zA-Z0-9]/-}-${run_idx}" + local log_file="$server_results/warmup_console_run${run_idx}.log" + local -a pipeline_status + local k6_exit tee_exit + + FAILURE_STAGE="warmup:$server:run$run_idx" + FAILURE_REASON="k6 warmup setup or execution failed" + FAILURE_EVIDENCE_LOG="${log_file#"$RESULTS_DIR"/}" + CURRENT_K6_CONTAINER="$container" + set +e + docker run --rm --name "$container" \ + --network host --cpus "$BENCHMARK_K6_CPUS" --memory "$BENCHMARK_K6_MEMORY" \ + --cpuset-cpus "$HOST_CPUSET" \ + --user "$(id -u):$(id -g)" \ + -v "$K6_RUNTIME_SCRIPT_DIR:/scripts:ro" \ + -v "$server_results:/results" \ + -e SERVER_URL="$server_url" \ + -e SERVER_NAME="$server" \ + -e MCP_PROTOCOL_VERSION="$MCP_PROTOCOL_VERSION" \ + -e EXPECTED_PROTOCOL_VERSION="$negotiated_protocol_version" \ + -e BENCHMARK_CONTRACT="$ELIGIBILITY_CONTRACT" \ + -e BENCHMARK_MODE=warmup \ + -e BENCHMARK_VUS="$BENCHMARK_VUS" \ + -e BENCHMARK_RAMP_DURATION="$BENCHMARK_RAMP_DURATION" \ + -e BENCHMARK_WARMUP_DURATION="$BENCHMARK_WARMUP_DURATION" \ + -e OUTPUT_PATH="/results/warmup_summary_run${run_idx}.json" \ + "$K6_IMAGE" run /scripts/benchmark.js \ + 2>&1 | tee "$log_file" + pipeline_status=("${PIPESTATUS[@]}") + set -e + CURRENT_K6_CONTAINER="" + k6_exit="${pipeline_status[0]}" + tee_exit="${pipeline_status[1]}" + if [[ "$k6_exit" -ne 0 || "$tee_exit" -ne 0 ]]; then + record_failure_from_log "warmup:$server:run$run_idx" "$log_file" \ + "k6 warmup failed correctness or load thresholds" + if [[ "$k6_exit" -ne 0 ]]; then + return "$k6_exit" + fi + return "$tee_exit" + fi + clear_failure_context +} - info "Warmup for $server_name: 3 full sessions per tool" - for _ in {1..3}; do - warmup_single_session "$server_url" "search_products" '{"category":"Electronics","min_price":50,"max_price":500,"limit":10}' - warmup_single_session "$server_url" "get_user_cart" '{"user_id":"user-00001"}' - warmup_single_session "$server_url" "checkout" '{"user_id":"user-00001","items":[{"product_id":42,"quantity":2},{"product_id":1337,"quantity":1}]}' +run_k6_measurement() { + local server="$1" + local server_url="$2" + local server_results="$3" + local run_idx="$4" + local negotiated_protocol_version="$5" + verify_runtime_harness + local container="mcp-k6-measure-${server//[^a-zA-Z0-9]/-}-${run_idx}" + local k6_pid k6_exit=0 collector_exit=0 + local log_file="$server_results/k6_console_run${run_idx}.log" + local resource_log="$server_results/resource_headroom_run${run_idx}.log" + local -a resource_pipeline_status + local resource_exit tee_exit + + FAILURE_STAGE="measurement:$server:run$run_idx" + FAILURE_REASON="k6 measurement setup or execution failed" + FAILURE_EVIDENCE_LOG="${log_file#"$RESULTS_DIR"/}" + CURRENT_K6_CONTAINER="$container" + docker rm -f "$container" >/dev/null 2>&1 || true + docker run --name "$container" \ + --network host --cpus "$BENCHMARK_K6_CPUS" --memory "$BENCHMARK_K6_MEMORY" \ + --cpuset-cpus "$HOST_CPUSET" \ + --user "$(id -u):$(id -g)" \ + -v "$K6_RUNTIME_SCRIPT_DIR:/scripts:ro" \ + -v "$server_results:/results" \ + -e SERVER_URL="$server_url" \ + -e SERVER_NAME="$server" \ + -e MCP_PROTOCOL_VERSION="$MCP_PROTOCOL_VERSION" \ + -e EXPECTED_PROTOCOL_VERSION="$negotiated_protocol_version" \ + -e BENCHMARK_CONTRACT="$ELIGIBILITY_CONTRACT" \ + -e BENCHMARK_MODE=measurement \ + -e BENCHMARK_VUS="$BENCHMARK_VUS" \ + -e BENCHMARK_MEASURE_DURATION="$BENCHMARK_MEASURE_DURATION" \ + -e OUTPUT_PATH="/results/k6_summary_run${run_idx}.json" \ + "$K6_IMAGE" run --paused --address 127.0.0.1:6565 /scripts/benchmark.js \ + 2>&1 | tee "$log_file" & + k6_pid=$! + + wait_for_container "$container" + mkdir -p "$server_results/k6" + python3 "$SCRIPT_DIR/capture_container.py" \ + --container "$container" --run "$run_idx" \ + --output-dir "$server_results/k6" \ + --expected-cpus "$BENCHMARK_K6_CPUS" \ + --expected-memory-bytes "$BENCHMARK_K6_MEMORY_BYTES" \ + --expected-cpuset "$HOST_CPUSET" \ + --skip-executable + + start_collector "${CONTAINER_NAME[$server]}" "$server_results/stats_run${run_idx}.json" + start_collector mcp-redis "$server_results/redis_stats_run${run_idx}.json" + start_collector mcp-api-service "$server_results/api_stats_run${run_idx}.json" + start_collector "$container" "$server_results/k6_stats_run${run_idx}.json" + start_collector @host "$server_results/host_stats_run${run_idx}.json" + if ! wait_for_collectors_ready 15; then + set +e + stop_collectors + docker rm -f "$container" >/dev/null 2>&1 + wait "$k6_pid" 2>/dev/null + set -e + CURRENT_K6_CONTAINER="" + record_resource_collection_failure \ + "resource_collection:$server:run$run_idx:readiness" \ + "$server_results" "$run_idx" \ + "resource collectors did not become ready" + error "[$server] resource collectors failed before run $run_idx" + return 1 + fi + resume_k6 "$server" "$run_idx" + + set +e + wait "$k6_pid" + k6_exit=$? + stop_collectors + collector_exit=$? + set -e + docker rm "$container" >/dev/null + CURRENT_K6_CONTAINER="" + + if [[ "$k6_exit" -ne 0 ]]; then + record_failure_from_log "measurement:$server:run$run_idx" "$log_file" \ + "k6 measurement failed correctness or load thresholds" + error "[$server] measured k6 run $run_idx failed correctness or load thresholds" + return "$k6_exit" + fi + if [[ "$collector_exit" -ne 0 ]]; then + record_resource_collection_failure \ + "resource_collection:$server:run$run_idx" \ + "$server_results" "$run_idx" "resource collector failed" + error "[$server] resource collector failed during run $run_idx" + return "$collector_exit" + fi + set +e + python3 "$SCRIPT_DIR/validate_resource_headroom.py" \ + "$server_results/resource_headroom_run${run_idx}.json" \ + --expected-duration "$BENCHMARK_MEASURE_DURATION" \ + --observed-resource "server:$server_results/stats_run${run_idx}.json:2:2147483648" \ + --resource "redis:$server_results/redis_stats_run${run_idx}.json:0.5:536870912" \ + --resource "api_service:$server_results/api_stats_run${run_idx}.json:2:2147483648" \ + --resource "load_generator:$server_results/k6_stats_run${run_idx}.json:4:2147483648" \ + --resource "host:$server_results/host_stats_run${run_idx}.json:$HOST_CPU_LIMIT:$HOST_MEMORY_BYTES" \ + 2>&1 | tee "$resource_log" + resource_pipeline_status=("${PIPESTATUS[@]}") + set -e + resource_exit="${resource_pipeline_status[0]}" + tee_exit="${resource_pipeline_status[1]}" + if [[ "$resource_exit" -ne 0 || "$tee_exit" -ne 0 ]]; then + record_failure_from_log "resource_headroom:$server:run$run_idx" "$resource_log" \ + "resource headroom validation failed" + if [[ "$resource_exit" -ne 0 ]]; then + return "$resource_exit" + fi + return "$tee_exit" + fi + clear_failure_context +} + +write_run_manifest() { + local expected_images_json + local supplemental_targets_json='[]' + local target + for target in "${selected_servers[@]}"; do + if [[ "${REQUIRE_SUPPLEMENTAL[$target]:-0}" -eq 1 ]]; then + supplemental_targets_json="$( + jq -c --arg target "$target" '. + [$target]' \ + <<< "$supplemental_targets_json" + )" + fi done + expected_images_json="$( + printf '%s\n' "${expected_image_containers[@]}" | jq -R . | jq -s . + )" + compose images --format json | jq -s --argjson expected "$expected_images_json" ' + (if length == 1 and (.[0] | type) == "array" then .[0] else . end) + | map(select(.ContainerName as $name | $expected | index($name))) + | . as $images + | ($images | map(.ContainerName)) as $actual + | if length > 0 + and all(.[]; + type == "object" + and (.ContainerName | type) == "string" + and (.ContainerName | length) > 0 + and (.ID | type) == "string" + and (.ID | startswith("sha256:")) + and (.ID | length) > 7) + and (($actual | length) == ($actual | unique | length)) + and (($actual | sort) == ($expected | sort)) + then $images + else error("compose images did not contain the exact selected container set with immutable IDs") + end + ' > "$RESULTS_DIR/compose_images.json" + jq -n \ + --arg started_at "$(date -u +%Y-%m-%dT%H:%M:%SZ)" \ + --arg profile "$RUN_PROFILE" \ + --argjson runs "$BENCHMARK_RUNS" \ + --argjson vus "$BENCHMARK_VUS" \ + --argjson request_concurrency "$BENCHMARK_REQUEST_CONCURRENCY" \ + --arg ramp "$BENCHMARK_RAMP_DURATION" \ + --arg warmup "$BENCHMARK_WARMUP_DURATION" \ + --arg measurement "$BENCHMARK_MEASURE_DURATION" \ + --arg k6_image "$K6_IMAGE" \ + --arg eligibility_contract "$ELIGIBILITY_CONTRACT" \ + --arg supplemental_contract "$SUPPLEMENTAL_CONTRACT" \ + --arg host_network_policy "$HOST_NETWORK_POLICY" \ + --argjson supplemental_targets "$supplemental_targets_json" \ + --argjson k6_cpus "$BENCHMARK_K6_CPUS" \ + --argjson k6_memory_bytes "$BENCHMARK_K6_MEMORY_BYTES" \ + --slurpfile schedule "$RESULTS_DIR/run_order.json" \ + --slurpfile images "$RESULTS_DIR/compose_images.json" \ + '{schema_version: 2, status: "running", started_at: $started_at, + profile: $profile, publishable_candidate: false, + parameters: {runs: $runs, vus: $vus, + request_concurrency: $request_concurrency, ramp_duration: $ramp, + warmup_duration: $warmup, measurement_duration: $measurement, + k6_cpus: $k6_cpus, k6_memory_bytes: $k6_memory_bytes}, + contracts: { + eligibility: $eligibility_contract, + supplemental: { + name: $supplemental_contract, + required_targets: $supplemental_targets + } + }, + resource_collection: {host_network_policy: $host_network_policy}, + k6_image: $k6_image, schedule: $schedule[0], images: $images[0]}' \ + > "$RESULTS_DIR/run_manifest.json" +} - ok "Warmup finished for $server_name" +generate_comparison() { + local comparison_file="$RESULTS_DIR/comparison.txt" + { + printf "Benchmark comparison (%s profile)\n" "$RUN_PROFILE" + printf "Only separate constant-VU measurement invocations are reported.\n" + printf "Requested MCP revision: %s; each row records the server-selected revision.\n" \ + "$MCP_PROTOCOL_VERSION" + printf "Eligibility contract: %s (identical measured predicates for every row).\n" \ + "$ELIGIBILITY_CONTRACT" + printf "Supplemental adapter contract: %s (out-of-band evidence).\n" \ + "$SUPPLEMENTAL_CONTRACT" + printf "Results: %s\n\n" "$RESULTS_DIR" + printf "%-18s %-12s %-12s %-10s %-8s %-12s %-12s %-12s %-12s\n" \ + "Server" "Protocol" "Operations" "Ops/s" "CV%" "p50(ms)" "p95(ms)" "p99(ms)" "ErrorRate" + printf "%-18s %-12s %-12s %-10s %-8s %-12s %-12s %-12s %-12s\n" \ + "------------------" "------------" "------------" "----------" "--------" \ + "------------" "------------" "------------" "------------" + local server summary multi_stats protocol operations rps p50 p95 p99 err_rate cv_pct + for server in "${selected_servers[@]}"; do + summary="$RESULTS_DIR/$server/k6_summary.json" + multi_stats="$RESULTS_DIR/$server/k6_multi_run_stats.json" + operations="$(jq -r '.rates.operations.count' "$summary")" + protocol="$(jq -r '.negotiated_protocol_version' "$multi_stats")" + rps="$(jq -r '.rates.operations.per_second' "$summary")" + p50="$(jq -r '.latency.combined_tool_call.p50_ms' "$summary")" + p95="$(jq -r '.latency.combined_tool_call.p95_ms' "$summary")" + p99="$(jq -r '.latency.combined_tool_call.p99_ms' "$summary")" + err_rate="$(jq -r '.errors.mcp_rate' "$summary")" + cv_pct="$(jq -r '.sample_cv_pct' "$multi_stats")" + printf "%-18s %-12s %-12s %-10.2f %-8.2f %-12.2f %-12.2f %-12.2f %-12.4f\n" \ + "$server" "$protocol" "$operations" "$rps" "$cv_pct" "$p50" "$p95" "$p99" "$err_rate" + done + } | tee "$comparison_file" } selected_servers=() if [[ "$#" -eq 0 ]]; then - selected_servers=(cpp python go rust) + selected_servers=("${CPP_SDK_SERVERS[@]}") +elif [[ "$#" -eq 1 && "$1" == "cpp-sdks" ]]; then + selected_servers=("${CPP_SDK_SERVERS[@]}") +elif [[ "$#" -eq 1 && "$1" == "baseline" ]]; then + selected_servers=("${BASELINE_SERVERS[@]}") elif [[ "$#" -eq 1 && "$1" == "all" ]]; then - selected_servers=(cpp python go rust) + selected_servers=("${CPP_SDK_SERVERS[@]}" "${BASELINE_SERVERS[@]}") elif [[ "$#" -eq 1 && "$1" == *","* ]]; then IFS=',' read -r -a selected_servers <<< "$1" else selected_servers=("$@") fi +load_alternative_manifest + +declare -A seen_servers=() for server in "${selected_servers[@]}"; do case "$server" in - cpp|python|go|rust) ;; + cpp|ours-comparable|python|go|rust|hkr04|fastmcpp|cxxmcp|neumann) ;; + gopher-mcp) + error "gopher-mcp exposes legacy HTTP+SSE rather than the shared /mcp Streamable HTTP contract" + exit 2 + ;; -h|--help) usage exit 0 ;; *) - error "Unknown server '$server'. Allowed: cpp, python, go, rust, all" + error "Unknown server '$server'" usage exit 1 ;; esac + [[ -z "${seen_servers[$server]:-}" ]] || { + error "Duplicate server selection: $server" + exit 1 + } + seen_servers[$server]=1 done -mkdir -p "$RESULTS_DIR" - -info "Step 1/6: Pre-flight checks" -docker compose version >/dev/null || { error "docker compose not found"; exit 1; } -command -v python3 >/dev/null || { error "python3 not found"; exit 1; } -command -v jq >/dev/null || { error "jq not found"; exit 1; } -ensure_upstream_benchmark_servers || exit 1 -ok "Pre-flight checks passed" - -info "Step 2/6: Start infrastructure" -docker compose -f "$SCRIPT_DIR/docker-compose.yml" up -d redis api-service -wait_for_http "api-service" "http://localhost:8100/health" 30 +[[ "${#selected_servers[@]}" -gt 0 ]] || { + error "At least one server must be selected" + exit 1 +} +[[ "$BENCHMARK_RUNS" =~ ^[1-9][0-9]*$ && $((BENCHMARK_RUNS % 2)) -eq 1 ]] || { + error "BENCHMARK_RUNS must be a positive odd integer" + exit 1 +} +[[ "$BENCHMARK_VUS" =~ ^[1-9][0-9]*$ ]] || { + error "BENCHMARK_VUS must be a positive integer" + exit 1 +} +[[ "$BENCHMARK_ORDER_SEED" =~ ^[0-9]+$ ]] || { + error "BENCHMARK_ORDER_SEED must be a non-negative integer" + exit 1 +} -info "Waiting for redis readiness" -start_redis_wait="$(date +%s)" -while true; do - if docker compose -f "$SCRIPT_DIR/docker-compose.yml" exec -T redis redis-cli ping >/dev/null 2>&1; then - ok "redis is healthy" - break - fi - if (( $(date +%s) - start_redis_wait >= 30 )); then - error "Timeout waiting for redis" - exit 1 - fi - sleep 1 +RUN_PROFILE="smoke" +baseline_selection=true +for server in "${selected_servers[@]}"; do + case "$server" in + python|go|rust) ;; + *) baseline_selection=false ;; + esac done +if [[ "$BENCHMARK_RUNS" -eq 3 \ + && "$BENCHMARK_VUS" -eq 50 \ + && "$BENCHMARK_RAMP_DURATION" == "15s" \ + && "$BENCHMARK_WARMUP_DURATION" == "60s" \ + && "$BENCHMARK_MEASURE_DURATION" == "5m" ]]; then + if [[ "${selected_servers[*]}" == "${CPP_SDK_SERVERS[*]}" ]]; then + RUN_PROFILE="production" + elif [[ "$baseline_selection" == true ]]; then + RUN_PROFILE="baseline-diagnostic" + fi +fi +if [[ "$RUN_PROFILE" == "production" \ + && "$BENCHMARK_VUS" -ne "$BENCHMARK_REQUEST_CONCURRENCY" ]]; then + error "Production VUs must match request concurrency ($BENCHMARK_REQUEST_CONCURRENCY)" + exit 1 +fi -info "Step 3/6: Seed Redis" -docker rm -f mcp-redis-seeder >/dev/null 2>&1 || true -docker compose -f "$SCRIPT_DIR/docker-compose.yml" --profile seeder run --rm redis-seeder -ok "Redis seeding completed" - -info "Step 4/6: Benchmark selected servers: ${selected_servers[*]}" -for server_name in "${selected_servers[@]}"; do - service_name="${SERVICE_NAME[$server_name]}" - container_name="${CONTAINER_NAME[$server_name]}" - mcp_url="${MCP_URL[$server_name]}" - health_url="${HEALTH_URL[$server_name]}" - server_results="$RESULTS_DIR/$server_name" - mkdir -p "$server_results" - - info "[$server_name] Reset Redis" - docker compose -f "$SCRIPT_DIR/docker-compose.yml" exec -T redis redis-cli FLUSHDB >/dev/null - docker rm -f mcp-redis-seeder >/dev/null 2>&1 || true - docker compose -f "$SCRIPT_DIR/docker-compose.yml" --profile seeder run --rm redis-seeder >/dev/null - ok "[$server_name] Redis reset/reseed done" +acquire_benchmark_lock - info "[$server_name] Stop all MCP servers" - docker compose -f "$SCRIPT_DIR/docker-compose.yml" stop cpp-server python-server go-server rust-server>/dev/null +timestamp="$(date +%Y%m%d_%H%M%S)" +RESULTS_DIR="${BENCHMARK_RESULTS_DIR:-$SCRIPT_DIR/results/${timestamp}_${RUN_PROFILE}}" +if [[ -e "$RESULTS_DIR" \ + && -n "$(find "$RESULTS_DIR" -mindepth 1 -maxdepth 1 -print -quit)" ]]; then + error "Results directory is not empty: $RESULTS_DIR" + exit 1 +fi +mkdir -p "$RESULTS_DIR" - info "[$server_name] Start target server: $service_name" - docker compose -f "$SCRIPT_DIR/docker-compose.yml" up -d "$service_name" +info "Preflight: corrected $RUN_PROFILE profile for ${selected_servers[*]}" +command -v docker >/dev/null || { error "docker not found"; exit 1; } +command -v git >/dev/null || { error "git not found"; exit 1; } +command -v jq >/dev/null || { error "jq not found"; exit 1; } +command -v python3 >/dev/null || { error "python3 not found"; exit 1; } +command -v curl >/dev/null || { error "curl not found"; exit 1; } +HOST_CPU_LIMIT="$(python3 -c 'import os; print(len(os.sched_getaffinity(0)))')" +HOST_CPUSET="$(python3 -c 'import os; print(",".join(str(cpu) for cpu in sorted(os.sched_getaffinity(0))))')" +HOST_MEMORY_BYTES="$(awk '/^MemTotal:/ {printf "%.0f", $2 * 1024; exit}' /proc/meminfo)" +[[ "$HOST_CPU_LIMIT" =~ ^[1-9][0-9]*$ ]] || { error "Could not determine host CPU affinity"; exit 1; } +[[ "$HOST_MEMORY_BYTES" =~ ^[1-9][0-9]*$ ]] || { error "Could not determine host memory"; exit 1; } +[[ -n "$HOST_CPUSET" ]] || { error "Could not determine host CPU set"; exit 1; } +export BENCHMARK_CPUSET="$HOST_CPUSET" +if [[ "$RUN_PROFILE" != "smoke" && "$HOST_CPU_LIMIT" -lt 9 ]]; then + error "Full-duration runs require at least 9 host CPUs for 8.5 CPUs of configured limits" + exit 1 +fi +compose version >/dev/null +[[ -f "$K6_SCRIPT" ]] || { error "Missing local k6 profile: $K6_SCRIPT"; exit 1; } +for helper in benchmark_order.py capture_container.py capture_environment.py \ + collect_stats.py select_median_run.py summarize_collector_failures.py \ + validate_resource_headroom.py verify_server.py; do + [[ -f "$SCRIPT_DIR/$helper" ]] || { error "Missing benchmark helper: $helper"; exit 1; } +done - health_timeout=30 - if [[ "$server_name" == "cpp" ]]; then - health_timeout=60 +mkdir -p "$RESULTS_DIR/harness" +cp "$COMPOSE_FILE" "$RESULTS_DIR/harness/docker-compose.yml" +cp "$K6_SCRIPT" "$RESULTS_DIR/harness/benchmark.js" +( + cd "$RESULTS_DIR/harness" + sha256sum benchmark.js docker-compose.yml > SHA256SUMS +) +RUNTIME_COMPOSE_FILE="$RESULTS_DIR/harness/docker-compose.yml" +K6_RUNTIME_SCRIPT_DIR="$RESULTS_DIR/harness" +RUNTIME_COMPOSE_SHA256="$(sha256sum "$RUNTIME_COMPOSE_FILE" | awk '{print $1}')" +RUNTIME_K6_SHA256="$(sha256sum "$K6_RUNTIME_SCRIPT_DIR/benchmark.js" | awk '{print $1}')" +chmod 0444 "$RUNTIME_COMPOSE_FILE" "$K6_RUNTIME_SCRIPT_DIR/benchmark.js" \ + "$K6_RUNTIME_SCRIPT_DIR/SHA256SUMS" + +ensure_upstream_sources +for server in "${selected_servers[@]}"; do + if [[ -n "${ALTERNATIVE_REPO[$server]:-}" ]]; then + [[ "${ALTERNATIVE_STATUS[$server]}" == "supported" ]] || { + error "$server is not runnable: ${ALTERNATIVE_STATUS[$server]}" + exit 2 + } + ensure_alternative_source "$server" fi - wait_for_http "$server_name server" "$health_url" "$health_timeout" - - run_warmup "$server_name" "$mcp_url" - - info "[$server_name] Start stats collector" - python3 "$SCRIPT_DIR/collect_stats.py" "$container_name" "$server_results/stats.json" 1.0 & - stats_pid=$! - - K6_RUNS=3 - info "[$server_name] Run k6 benchmark ($K6_RUNS iterations)" - for run_idx in $(seq 1 "$K6_RUNS"); do - info "[$server_name] k6 run $run_idx/$K6_RUNS" +done - if (( run_idx > 1 )); then - info "[$server_name] Re-seed Redis before run $run_idx" - docker compose -f "$SCRIPT_DIR/docker-compose.yml" exec -T redis redis-cli FLUSHDB >/dev/null - docker rm -f mcp-redis-seeder >/dev/null 2>&1 || true - docker compose -f "$SCRIPT_DIR/docker-compose.yml" --profile seeder run --rm redis-seeder >/dev/null - run_warmup "$server_name" "$mcp_url" - fi +python3 "$SCRIPT_DIR/benchmark_order.py" \ + --seed "$BENCHMARK_ORDER_SEED" --runs "$BENCHMARK_RUNS" \ + "${selected_servers[@]}" > "$RESULTS_DIR/run_order.json" +compose config > "$RESULTS_DIR/compose.resolved.yml" + +info "Pulling pinned Redis and k6 images" +compose --profile seeder pull redis +docker pull "$K6_IMAGE" + +environment_args=( + "$RESULTS_DIR/environment.json" + --project-dir "$PROJECT_DIR" + --upstream-dir "$UPSTREAM_BENCHMARK_DIR" + --upstream-url "$UPSTREAM_BENCHMARK_REPO" + --alternative-root "$ALTERNATIVE_SOURCE_DIR" + --sources-manifest "$ALTERNATIVE_MANIFEST" + --protocol-version "$MCP_PROTOCOL_VERSION" + --eligibility-contract "$ELIGIBILITY_CONTRACT" + --k6-image "$K6_IMAGE" + --upstream-commit "$UPSTREAM_BENCHMARK_COMMIT" + --order-seed "$BENCHMARK_ORDER_SEED" + --runs "$BENCHMARK_RUNS" + --vus "$BENCHMARK_VUS" + --measurement-duration "$BENCHMARK_MEASURE_DURATION" + --warmup-duration "$BENCHMARK_WARMUP_DURATION" + --ramp-duration "$BENCHMARK_RAMP_DURATION" + --servers "${selected_servers[@]}" +) +python3 "$SCRIPT_DIR/capture_environment.py" "${environment_args[@]}" - docker run --rm \ - --network host \ - --user "$(id -u):$(id -g)" \ - -v "$UPSTREAM_BENCHMARK_DIR/benchmark:/scripts:ro" \ - -v "$server_results:/results" \ - -e SERVER_URL="$mcp_url" \ - -e SERVER_NAME="$server_name" \ - -e OUTPUT_PATH="/results/k6_summary_run${run_idx}.json" \ - grafana/k6:latest run /scripts/benchmark.js \ - 2>&1 | tee "$server_results/k6_console_run${run_idx}.log" - done +build_services=(api-service redis-seeder) +for server in "${selected_servers[@]}"; do + build_services+=("${SERVICE_NAME[$server]}") +done +info "Building every selected server once before the schedule" +compose build "${build_services[@]}" 2>&1 | tee "$RESULTS_DIR/build.log" +python3 "$SCRIPT_DIR/capture_environment.py" \ + --verify-source "$RESULTS_DIR/environment.json" --project-dir "$PROJECT_DIR" + +# Materialize stopped containers so Compose can report the exact built image +# IDs before execution. Every measured target is still force-recreated later. +image_services=(redis api-service) +expected_image_containers=(mcp-redis mcp-api-service mcp-redis-seeder) +for server in "${selected_servers[@]}"; do + image_services+=("${SERVICE_NAME[$server]}") + expected_image_containers+=("${CONTAINER_NAME[$server]}") +done +compose --profile seeder create "${image_services[@]}" redis-seeder >/dev/null - # Pick the median run by RPS and select as the canonical k6_summary.json - python3 "$SCRIPT_DIR/select_median_run.py" "$server_results" "$K6_RUNS" +write_run_manifest - info "[$server_name] Stop stats collector" - if kill -0 "$stats_pid" 2>/dev/null; then - kill "$stats_pid" 2>/dev/null || true - wait "$stats_pid" 2>/dev/null || true +info "Starting shared Redis and API service" +BENCHMARK_ACTIVE=1 +compose up -d --force-recreate redis api-service +wait_for_http "api-service" "http://localhost:8100/health" 60 +start_redis_wait="$(date +%s)" +until compose exec -T redis redis-cli ping >/dev/null 2>&1; do + if (( $(date +%s) - start_redis_wait >= 60 )); then + error "Timeout waiting for Redis" + exit 1 fi - stats_pid="" - - ok "[$server_name] Benchmark complete ($K6_RUNS runs, median selected)" + sleep 1 done -info "Step 5/6: Generate comparison summary" -comparison_file="$RESULTS_DIR/comparison.txt" - -{ - printf "Benchmark comparison\n" - printf "Results: %s\n\n" "$RESULTS_DIR" - printf "%-10s %-12s %-10s %-8s %-12s %-12s %-12s %-12s\n" "Server" "Requests" "RPS" "CV%" "p50(ms)" "p95(ms)" "p99(ms)" "ErrorRate" - printf "%-10s %-12s %-10s %-8s %-12s %-12s %-12s %-12s\n" "----------" "------------" "----------" "--------" "------------" "------------" "------------" "------------" - - for server_name in "${selected_servers[@]}"; do - summary="$RESULTS_DIR/$server_name/k6_summary.json" - multi_stats="$RESULTS_DIR/$server_name/k6_multi_run_stats.json" - if [[ ! -f "$summary" ]]; then - printf "%-10s %-12s %-10s %-8s %-12s %-12s %-12s %-12s\n" "$server_name" "N/A" "N/A" "N/A" "N/A" "N/A" "N/A" "N/A" - continue +mkdir -p "$RESULTS_DIR/shared/redis" "$RESULTS_DIR/shared/api_service" +python3 "$SCRIPT_DIR/capture_container.py" \ + --container mcp-redis --run 1 --output-dir "$RESULTS_DIR/shared/redis" \ + --expected-cpus 0.5 --expected-memory-bytes 536870912 \ + --expected-cpuset "$HOST_CPUSET" +python3 "$SCRIPT_DIR/capture_container.py" \ + --container mcp-api-service --run 1 \ + --output-dir "$RESULTS_DIR/shared/api_service" \ + --expected-cpus 2 --expected-memory-bytes 2147483648 \ + --expected-cpuset "$HOST_CPUSET" + +sequence_index=0 +for run_idx in $(seq 1 "$BENCHMARK_RUNS"); do + mapfile -t round_servers < <( + python3 "$SCRIPT_DIR/benchmark_order.py" \ + --seed "$BENCHMARK_ORDER_SEED" --runs "$BENCHMARK_RUNS" \ + --run "$run_idx" "${selected_servers[@]}" + ) + for server in "${round_servers[@]}"; do + sequence_index=$((sequence_index + 1)) + service="${SERVICE_NAME[$server]}" + container="${CONTAINER_NAME[$server]}" + server_url="${MCP_URL[$server]}" + server_results="$RESULTS_DIR/$server" + mkdir -p "$server_results" + CURRENT_SERVER="$container" + CURRENT_SERVER_RESULTS="$server_results" + + info "[$sequence_index] $server round $run_idx: recreate target" + compose stop "${ALL_MCP_SERVICES[@]}" >/dev/null + compose up -d --force-recreate --no-deps "$service" + if [[ "${HEALTH_KIND[$server]}" == "mcp" ]]; then + wait_for_mcp "$server" "${HEALTH_URL[$server]}" 90 + else + wait_for_http "$server" "${HEALTH_URL[$server]}" 90 fi - requests="$(jq -r '.mcp.total_mcp_requests // 0' "$summary")" - rps="$(jq -r '.http.rps // 0' "$summary")" - p50="$(jq -r '.http.latency.p50 // 0' "$summary")" - p95="$(jq -r '.http.latency.p95 // 0' "$summary")" - p99="$(jq -r '.http.latency.p99 // 0' "$summary")" - err_rate="$(jq -r '.mcp.error_rate // 0' "$summary")" - - cv_pct="N/A" - if [[ -f "$multi_stats" ]]; then - cv_pct="$(jq -r '.cv_pct' "$multi_stats")" + python3 "$SCRIPT_DIR/capture_container.py" \ + --container "$container" --run "$run_idx" --output-dir "$server_results" \ + --expected-cpus 2 --expected-memory-bytes 2147483648 \ + --expected-cpuset "$HOST_CPUSET" + + reset_redis_dataset + info "[$server] correctness and negotiation gate" + verifier_contract_args=( + --eligibility-contract "$ELIGIBILITY_CONTRACT" + --redis-url "redis://127.0.0.1:6379/0" + ) + if [[ "${REQUIRE_SUPPLEMENTAL[$server]:-0}" -eq 1 ]]; then + verifier_contract_args+=(--require-supplemental) fi - - printf "%-10s %-12s %-10.2f %-8s %-12.2f %-12.2f %-12.2f %-12.4f\n" \ - "$server_name" "$requests" "$rps" "${cv_pct}%" "$p50" "$p95" "$p99" "$err_rate" + protocol_preflight="$server_results/protocol_preflight_run${run_idx}.json" + preflight_log="$server_results/protocol_preflight_run${run_idx}.log" + run_protocol_verifier "protocol_preflight:$server:run$run_idx" \ + "$preflight_log" "$server_url" --name "$server" \ + --expected-server-type "${EXPECTED_SERVER_TYPE[$server]}" \ + "${verifier_contract_args[@]}" \ + --output "$protocol_preflight" + negotiated_protocol_version="$( + jq -er '.negotiated_protocol_version | select(type == "string" and length > 0)' \ + "$protocol_preflight" + )" + reset_redis_dataset + + info "[$server] separate warmup (round $run_idx)" + run_k6_warmup "$server" "$server_url" "$server_results" "$run_idx" \ + "$negotiated_protocol_version" + + # Warmup mutates checkout history and rate-limit state. Keep the process + # alive for warm runtime caches, but restore the canonical dataset. + reset_redis_dataset + + info "[$server] measured constant-VU run $run_idx/$BENCHMARK_RUNS" + run_k6_measurement "$server" "$server_url" "$server_results" "$run_idx" \ + "$negotiated_protocol_version" + + # Re-seed before the postflight so the universal workload contract and + # supplemental diagnostics are evaluated from the same canonical state. + reset_redis_dataset + postflight_log="$server_results/protocol_postflight_run${run_idx}.log" + run_protocol_verifier "protocol_postflight:$server:run$run_idx" \ + "$postflight_log" "$server_url" \ + --name "$server-post-run" \ + --expected-protocol-version "$negotiated_protocol_version" \ + --expected-server-type "${EXPECTED_SERVER_TYPE[$server]}" \ + "${verifier_contract_args[@]}" \ + --output "$server_results/protocol_postflight_run${run_idx}.json" + docker logs "$container" > "$server_results/server_run${run_idx}.log" 2>&1 + compose stop "$service" >/dev/null + CURRENT_SERVER="" + CURRENT_SERVER_RESULTS="" + ok "[$server] round $run_idx passed" done -} | tee "$comparison_file" +done -ok "Comparison saved to $comparison_file" +python3 "$SCRIPT_DIR/capture_environment.py" \ + --verify-source "$RESULTS_DIR/environment.json" --project-dir "$PROJECT_DIR" +for server in "${selected_servers[@]}"; do + python3 "$SCRIPT_DIR/select_median_run.py" "$RESULTS_DIR/$server" "$BENCHMARK_RUNS" +done -info "Step 6/6: Cleanup" -ok "Benchmark finished. Results directory: $RESULTS_DIR" -warn "Docker services left running intentionally for inspection" +generate_comparison +verify_runtime_harness +python3 "$SCRIPT_DIR/capture_environment.py" \ + --verify-source "$RESULTS_DIR/environment.json" --project-dir "$PROJECT_DIR" +RUN_SUCCEEDED=1 +ok "Corrected $RUN_PROFILE benchmark complete: $RESULTS_DIR" diff --git a/benchmark/select_median_run.py b/benchmark/select_median_run.py old mode 100755 new mode 100644 index 87248f1..dc659b5 --- a/benchmark/select_median_run.py +++ b/benchmark/select_median_run.py @@ -1,70 +1,456 @@ #!/usr/bin/env python3 -""" -Select median k6 benchmark run by RPS and generate multi-run statistics. +"""Select the median measured run and pair it with matching resource evidence.""" -Usage: - select_median_run.py - -Reads k6_summary_run{1..N}.json files, selects the median by RPS, copies it -to k6_summary.json, and writes k6_multi_run_stats.json with CV% and per-run stats. -""" +from __future__ import annotations import json -import sys -import os import math import shutil +import statistics +import sys +from pathlib import Path +from typing import Any + + +ELIGIBILITY_CONTRACT = "upstream-v2-strict-mcp-v1" +HOST_NETWORK_POLICY = "baseline-cohort-retire-v1" +OPERATION_SHAPE_CHECKS = ( + "search_products returns a valid tool result", + "get_user_cart returns a valid tool result", + "checkout returns a valid tool result", + "tools/list returns a valid tool collection", +) +OPERATION_CONTRACT_CHECKS = ( + "search_products satisfies the benchmark contract", + "get_user_cart satisfies the benchmark contract", + "checkout satisfies the benchmark contract", + "tools/list satisfies the benchmark contract", +) + + +def finite_number(value: Any, description: str) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError(f"{description} is not numeric") + number = float(value) + if not math.isfinite(number): + raise ValueError(f"{description} is not finite") + return number + + +def exact_nonnegative_integer(value: Any, description: str) -> int: + number = finite_number(value, description) + if number != int(number): + raise ValueError(f"{description} is not an exact integer") + parsed = int(number) + if parsed < 0: + raise ValueError(f"{description} is negative") + return parsed + + +def operation_rps(summary: dict[str, Any]) -> float: + rates = summary.get("rates", {}) + operations = rates.get("operations", {}) + if isinstance(operations, dict) and "per_second" in operations: + return float(operations["per_second"]) + return float(summary.get("http", {}).get("rps", 0)) + + +def validate_measurement_summary(summary: dict[str, Any], path: Path) -> None: + if not isinstance(summary, dict): + raise ValueError(f"{path} is not a JSON object") + config = summary.get("config") + if not isinstance(config, dict) or config.get("mode") != "measurement": + raise ValueError(f"{path} is not a measurement summary") + if config.get("eligibility_contract") != ELIGIBILITY_CONTRACT: + raise ValueError(f"{path} used a different eligibility contract") + rates = summary.get("rates") + operations = rates.get("operations") if isinstance(rates, dict) else None + if not isinstance(operations, dict) or "per_second" not in operations: + raise ValueError(f"{path} has no complete benchmark-operation rate") + rps = finite_number(operations["per_second"], f"{path} operation rate") + if "count" not in operations: + raise ValueError(f"{path} has no benchmark-operation count") + operation_count = exact_nonnegative_integer( + operations["count"], f"{path} operation count" + ) + if operation_count <= 0: + raise ValueError(f"{path} has no measured benchmark operations") + if rps <= 0: + raise ValueError(f"{path} has no measured benchmark operations") + + errors = summary.get("errors") + required_error_fields = { + "mcp", + "http", + "checks", + "mcp_rate", + "http_rate", + "check_pass_rate", + } + if not isinstance(errors, dict) or not required_error_fields.issubset(errors): + raise ValueError(f"{path} has incomplete correctness metrics") + values = { + name: finite_number(errors[name], f"{path} errors.{name}") + for name in required_error_fields + } + if ( + values["mcp"] != 0 + or values["http"] != 0 + or values["checks"] != 0 + or values["mcp_rate"] != 0 + or values["http_rate"] != 0 + or values["check_pass_rate"] != 1 + ): + raise ValueError(f"{path} failed the correctness gate") + + check_breakdown = summary.get("check_breakdown") + if not isinstance(check_breakdown, dict): + raise ValueError(f"{path} has no per-operation correctness breakdown") + for check_group in (OPERATION_SHAPE_CHECKS, OPERATION_CONTRACT_CHECKS): + validated_operations = 0 + for check_name in check_group: + check_counts = check_breakdown.get(check_name) + if not isinstance(check_counts, dict): + raise ValueError(f"{path} is missing correctness check {check_name!r}") + passes = exact_nonnegative_integer( + check_counts.get("passes"), f"{path} {check_name!r} passes" + ) + fails = exact_nonnegative_integer( + check_counts.get("fails"), f"{path} {check_name!r} failures" + ) + if fails != 0: + raise ValueError(f"{path} failed correctness check {check_name!r}") + validated_operations += passes + if operation_count > validated_operations: + raise ValueError( + f"{path} counted {operation_count} operations but only " + f"{validated_operations} reached the required checks" + ) + + +def load_json(path: Path) -> Any: + with path.open(encoding="utf-8") as stream: + return json.load(stream) + + +def write_json(path: Path, value: Any) -> None: + with path.open("w", encoding="utf-8") as stream: + json.dump(value, stream, indent=2, allow_nan=False) + stream.write("\n") + + +def interface_set(value: Any, description: str) -> set[str]: + if ( + not isinstance(value, list) + or any(not isinstance(item, str) or not item for item in value) + or len(value) != len(set(value)) + ): + raise ValueError(f"{description} is not a unique string list") + return set(value) + +def host_network_observation( + observed_bytes: dict[str, int], audit: dict[str, Any] +) -> dict[str, Any]: + host = audit.get("host") + if not isinstance(host, dict): + raise ValueError("host collector audit has no host metadata") + if host.get("network_interface_policy") != HOST_NETWORK_POLICY: + raise ValueError("host collector audit used an unknown network policy") -def main(): + initial = interface_set(host.get("network_interfaces"), "initial interfaces") + active = interface_set( + host.get("active_network_interfaces"), "active interfaces" + ) + retired = interface_set( + host.get("retired_network_interfaces"), "retired interfaces" + ) + ignored_new = interface_set( + host.get("ignored_new_network_interfaces"), "ignored new interfaces" + ) + if not initial or active & retired or active | retired != initial: + raise ValueError("host collector audit has an inconsistent baseline cohort") + if ignored_new & initial: + raise ValueError("host collector audit classifies baseline interfaces as new") + + return { + "policy": HOST_NETWORK_POLICY, + "scope": "baseline_interface_cohort", + "coverage": "partial" if retired or ignored_new else "complete", + "observed_bytes_during_collection": observed_bytes, + "initial_interfaces": sorted(initial), + "active_interfaces_at_end": sorted(active), + "retired_interfaces": sorted(retired), + "ignored_new_interfaces": sorted(ignored_new), + } + + +def resource_summary( + samples: list[dict[str, Any]], + selected_run: int, + *, + host_audit: dict[str, Any] | None = None, +) -> dict[str, Any]: + if not samples: + raise ValueError("selected resource sample file is empty") + + cpu_values: list[float] = [] + memory_values: list[int] = [] + memory_limits: list[int] = [] + network_values: list[tuple[int, int]] = [] + for index, sample in enumerate(samples): + if not isinstance(sample, dict): + raise ValueError(f"resource sample {index} is not an object") + required = { + "cpu_percent", + "mem_usage_bytes", + "mem_limit_bytes", + "net_io_rx", + "net_io_tx", + } + if not required.issubset(sample): + raise ValueError(f"resource sample {index} is incomplete") + cpu = finite_number(sample["cpu_percent"], f"sample {index} CPU") + memory = exact_nonnegative_integer( + sample["mem_usage_bytes"], f"sample {index} memory" + ) + memory_limit = exact_nonnegative_integer( + sample["mem_limit_bytes"], f"sample {index} memory limit" + ) + net_rx = exact_nonnegative_integer( + sample["net_io_rx"], f"sample {index} received bytes" + ) + net_tx = exact_nonnegative_integer( + sample["net_io_tx"], f"sample {index} transmitted bytes" + ) + if cpu < 0 or memory_limit <= 0: + raise ValueError(f"resource sample {index} contains an invalid counter or limit") + cpu_values.append(cpu) + memory_values.append(memory) + memory_limits.append(memory_limit) + network_values.append((net_rx, net_tx)) + if any( + later_rx < earlier_rx or later_tx < earlier_tx + for (earlier_rx, earlier_tx), (later_rx, later_tx) in zip( + network_values, network_values[1:] + ) + ): + raise ValueError("resource network counters are not monotonic") + first = samples[0] + last = samples[-1] + observed_network_bytes = { + "rx": network_values[-1][0] - network_values[0][0], + "tx": network_values[-1][1] - network_values[0][1], + } + summary = { + "selected_run": selected_run, + "sample_count": len(samples), + "first_timestamp": first.get("timestamp"), + "last_timestamp": last.get("timestamp"), + "cpu_percent": { + "mean": statistics.fmean(cpu_values), + "median": statistics.median(cpu_values), + "max": max(cpu_values), + }, + "memory_bytes": { + "mean": statistics.fmean(memory_values), + "median": statistics.median(memory_values), + "max": max(memory_values), + "limit": max(memory_limits), + }, + } + if host_audit is None: + summary["network_bytes_during_collection"] = observed_network_bytes + else: + summary["network_observation"] = host_network_observation( + observed_network_bytes, host_audit + ) + return summary + + +def require_file(path: Path, description: str) -> None: + if not path.is_file(): + raise ValueError(f"missing {description}: {path}") + + +def require_valid_report(path: Path, description: str) -> dict[str, Any]: + require_file(path, description) + report = load_json(path) + if not isinstance(report, dict) or report.get("valid") is not True: + raise ValueError(f"{description} is not valid: {path}") + return report + + +def main() -> int: if len(sys.argv) != 3: print(f"Usage: {sys.argv[0]} ", file=sys.stderr) - sys.exit(1) - - results_dir = sys.argv[1] - num_runs = int(sys.argv[2]) - - # Load all runs and extract RPS - runs = [] - for i in range(1, num_runs + 1): - path = os.path.join(results_dir, f'k6_summary_run{i}.json') - with open(path) as f: - data = json.load(f) - rps = data.get('http', {}).get('rps', 0) - runs.append((rps, i, path)) - - # Sort by RPS and select median - runs.sort(key=lambda x: x[0]) - median_idx = len(runs) // 2 - median_rps, median_run, median_path = runs[median_idx] - - # Compute coefficient of variation (CV%) - rps_values = [r[0] for r in runs] - mean_rps = sum(rps_values) / len(rps_values) - if mean_rps > 0: - variance = sum((v - mean_rps) ** 2 for v in rps_values) / len(rps_values) - std_dev = math.sqrt(variance) - cv_pct = (std_dev / mean_rps) * 100 - else: - cv_pct = 0.0 - - # Copy median run as canonical summary - canonical = os.path.join(results_dir, 'k6_summary.json') - shutil.copy2(median_path, canonical) - - # Write per-run stats and CV% - stats = { - 'runs': [{'run': r[1], 'rps': r[0]} for r in sorted(runs, key=lambda x: x[1])], - 'median_run': median_run, - 'median_rps': median_rps, - 'mean_rps': mean_rps, - 'cv_pct': round(cv_pct, 2), + return 1 + + results_dir = Path(sys.argv[1]) + try: + num_runs = int(sys.argv[2]) + except ValueError: + print("num_runs must be a positive odd number", file=sys.stderr) + return 1 + if num_runs < 1 or num_runs % 2 == 0: + print("num_runs must be a positive odd number", file=sys.stderr) + return 1 + + runs: list[dict[str, Any]] = [] + negotiated_versions: set[str] = set() + for run_number in range(1, num_runs + 1): + path = results_dir / f"k6_summary_run{run_number}.json" + summary = load_json(path) + validate_measurement_summary(summary, path) + preflight_path = results_dir / f"protocol_preflight_run{run_number}.json" + postflight_path = results_dir / f"protocol_postflight_run{run_number}.json" + preflight = load_json(preflight_path) + postflight = load_json(postflight_path) + if not isinstance(preflight, dict) or not isinstance(postflight, dict): + raise ValueError(f"run {run_number} protocol evidence is not an object") + negotiated_version = preflight.get("negotiated_protocol_version") + if ( + not isinstance(negotiated_version, str) + or not negotiated_version + or preflight.get("eligibility_contract") != ELIGIBILITY_CONTRACT + or postflight.get("eligibility_contract") != ELIGIBILITY_CONTRACT + or preflight.get("eligibility_valid") is not True + or postflight.get("eligibility_valid") is not True + or postflight.get("negotiated_protocol_version") != negotiated_version + or summary.get("config", {}).get( + "expected_negotiated_protocol_version" + ) + != negotiated_version + or summary.get("config", {}).get("eligibility_contract") + != ELIGIBILITY_CONTRACT + ): + raise ValueError(f"run {run_number} has inconsistent protocol evidence") + for phase, evidence in (("preflight", preflight), ("postflight", postflight)): + supplemental = evidence.get("supplemental_validation") + if not isinstance(supplemental, dict): + raise ValueError(f"run {run_number} {phase} has no supplemental evidence") + if supplemental.get("required") is True and supplemental.get("valid") is not True: + raise ValueError( + f"run {run_number} {phase} failed required supplemental validation" + ) + negotiated_versions.add(negotiated_version) + require_valid_report( + results_dir / f"resource_headroom_run{run_number}.json", + f"run {run_number} resource report", + ) + for cgroup_path, description in ( + (results_dir / f"cgroup_run{run_number}.json", "server cgroup evidence"), + (results_dir / "k6" / f"cgroup_run{run_number}.json", "k6 cgroup evidence"), + ): + require_file(cgroup_path, description) + cgroup = load_json(cgroup_path) + if ( + not isinstance(cgroup, dict) + or cgroup.get("validation", {}).get("ok") is not True + ): + raise ValueError(f"invalid {description}: {cgroup_path}") + runs.append( + { + "run": run_number, + "rps": operation_rps(summary), + "summary_path": path, + "protocol_version": negotiated_version, + } + ) + if len(negotiated_versions) != 1: + raise ValueError( + f"protocol negotiation changed between runs: {sorted(negotiated_versions)}" + ) + + ranked_runs = sorted(runs, key=lambda run: run["rps"]) + selected = ranked_runs[len(ranked_runs) // 2] + rps_values = [float(run["rps"]) for run in runs] + mean_rps = statistics.fmean(rps_values) + population_cv = statistics.pstdev(rps_values) / mean_rps * 100 if mean_rps else 0.0 + sample_cv = ( + statistics.stdev(rps_values) / mean_rps * 100 + if mean_rps and len(rps_values) > 1 + else 0.0 + ) + sample_standard_deviation = statistics.stdev(rps_values) if len(rps_values) > 1 else 0.0 + # Three runs are the production default. This exact t critical value is for df=2; + # omit the interval for other run counts rather than imply false precision. + confidence_interval: dict[str, float] | None = None + if len(rps_values) == 3: + margin = 4.302652729911275 * sample_standard_deviation / math.sqrt(3) + confidence_interval = {"low": mean_rps - margin, "high": mean_rps + margin} + + selected_run = int(selected["run"]) + shutil.copy2(selected["summary_path"], results_dir / "k6_summary.json") + + resource_files = { + "server": ("stats", "stats.json"), + "redis": ("redis_stats", "redis_stats.json"), + "api_service": ("api_stats", "api_stats.json"), + "load_generator": ("k6_stats", "k6_stats.json"), + "host": ("host_stats", "host_stats.json"), + } + resources: dict[str, Any] = {"selected_run": selected_run} + for resource_name, (run_prefix, canonical_name) in resource_files.items(): + selected_stats = results_dir / f"{run_prefix}_run{selected_run}.json" + require_file(selected_stats, f"selected {resource_name} samples") + shutil.copy2(selected_stats, results_dir / canonical_name) + audit_source = results_dir / f"{run_prefix}_run{selected_run}.audit.json" + audit_destination = results_dir / canonical_name.replace(".json", ".audit.json") + require_file(audit_source, f"selected {resource_name} collector audit") + shutil.copy2(audit_source, audit_destination) + audit = load_json(audit_source) + if not isinstance(audit, dict): + raise ValueError(f"selected {resource_name} collector audit is not an object") + samples = load_json(selected_stats) + if not isinstance(samples, list): + raise ValueError(f"selected {resource_name} samples are not a list") + resources[resource_name] = resource_summary( + samples, + selected_run, + host_audit=audit if resource_name == "host" else None, + ) + write_json(results_dir / "resource_summary.json", resources) + + selected_files = { + f"container_inspect_run{selected_run}.json": "container_inspect.json", + f"image_inspect_run{selected_run}.json": "image_inspect.json", + f"cgroup_run{selected_run}.json": "cgroup.json", + f"resource_headroom_run{selected_run}.json": "resource_headroom.json", + f"protocol_preflight_run{selected_run}.json": "protocol_preflight.json", + f"protocol_postflight_run{selected_run}.json": "protocol_postflight.json", } - with open(os.path.join(results_dir, 'k6_multi_run_stats.json'), 'w') as f: - json.dump(stats, f, indent=2) + for source_name, destination_name in selected_files.items(): + source = results_dir / source_name + require_file(source, f"selected-run evidence {source_name}") + shutil.copy2(source, results_dir / destination_name) + + selected_k6_dir = results_dir / "k6" + selected_k6_dir.mkdir(exist_ok=True) + for prefix in ("container_inspect", "image_inspect", "cgroup"): + source = selected_k6_dir / f"{prefix}_run{selected_run}.json" + require_file(source, f"selected k6 {prefix}") + shutil.copy2(source, selected_k6_dir / f"{prefix}.json") - print(f'Median run: {median_run} (RPS={median_rps:.2f}), CV%={cv_pct:.2f}%') + statistics_output = { + "runs": [{"run": run["run"], "rps": run["rps"]} for run in runs], + "median_run": selected_run, + "median_rps": selected["rps"], + "mean_rps": mean_rps, + "population_cv_pct": population_cv, + "sample_cv_pct": sample_cv, + "cv_pct": sample_cv, + "mean_rps_95pct_confidence_interval": confidence_interval, + "negotiated_protocol_version": next(iter(negotiated_versions)), + "eligibility_contract": ELIGIBILITY_CONTRACT, + } + write_json(results_dir / "k6_multi_run_stats.json", statistics_output) + print( + f"Median run: {selected_run} (RPS={selected['rps']:.2f}), " + f"sample CV={sample_cv:.2f}%" + ) + return 0 -if __name__ == '__main__': - main() +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/summarize_collector_failures.py b/benchmark/summarize_collector_failures.py new file mode 100644 index 0000000..581852b --- /dev/null +++ b/benchmark/summarize_collector_failures.py @@ -0,0 +1,69 @@ +#!/usr/bin/env python3 +"""Summarize failed resource-collector audits for benchmark manifests.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +from typing import Any + + +def _failure_detail(audit: dict[str, Any]) -> str: + fatal_error = audit.get("fatal_error") + if isinstance(fatal_error, str) and fatal_error: + return fatal_error + + failures = audit.get("failures") + if isinstance(failures, list): + for failure in reversed(failures): + if isinstance(failure, dict): + error = failure.get("error") + if isinstance(error, str) and error: + return error + + termination = audit.get("termination") + if isinstance(termination, dict): + reason = termination.get("reason") + if isinstance(reason, str) and reason: + return f"collector terminated: {reason}" + return "collector audit did not include an error" + + +def summarize_failures(audit_paths: list[Path]) -> str: + """Return one stable line describing every non-complete collector audit.""" + reasons: list[str] = [] + for audit_path in audit_paths: + try: + audit = json.loads(audit_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as error: + reasons.append(f"{audit_path.name}: unavailable audit ({error})") + continue + + if not isinstance(audit, dict): + reasons.append(f"{audit_path.name}: collector audit is not an object") + continue + if audit.get("status") == "complete": + continue + + target = audit.get("target") + if not isinstance(target, str) or not target: + target = audit_path.name + reasons.append(f"{target}: {_failure_detail(audit)}") + return "; ".join(reasons) + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("audits", nargs="+", type=Path) + args = parser.parse_args() + + summary = summarize_failures(args.audits) + if not summary: + return 1 + print(summary) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/test_benchmark_helpers.py b/benchmark/test_benchmark_helpers.py new file mode 100644 index 0000000..73a8fdf --- /dev/null +++ b/benchmark/test_benchmark_helpers.py @@ -0,0 +1,1212 @@ +#!/usr/bin/env python3 +"""Focused regression tests for the benchmark orchestration helpers.""" + +from __future__ import annotations + +import json +import signal +import sys +import tempfile +import unittest +from datetime import datetime, timedelta, timezone +from pathlib import Path +from unittest import mock + +from benchmark import benchmark_order +from benchmark import capture_container +from benchmark import capture_environment +from benchmark import collect_stats +from benchmark import select_median_run +from benchmark import summarize_collector_failures +from benchmark import validate_resource_headroom +from benchmark import verify_server + + +class CounterbalancedOrderTests(unittest.TestCase): + def test_five_servers_use_the_three_round_counterbalanced_design(self) -> None: + servers = ["ours", "hkr04", "fastmcpp", "cxxmcp", "neumann"] + orders = benchmark_order.counterbalanced_orders(servers, runs=3, seed=417) + + self.assertEqual(len(orders), 3) + base = orders[0] + self.assertEqual(orders[1], base[1:] + base[:1]) + self.assertEqual( + orders[2], + [base[4], base[3], base[0], base[2], base[1]], + ) + for order in orders: + self.assertEqual(set(order), set(servers)) + self.assertEqual(len(order), len(set(order))) + for server in servers: + positions = [order.index(server) for order in orders] + self.assertEqual(len(positions), len(set(positions))) + + +class MedianRunTests(unittest.TestCase): + def test_operation_rps_prefers_the_measured_operation_rate(self) -> None: + summary = { + "rates": {"operations": {"per_second": 321.5}}, + "http": {"rps": 999.0}, + } + self.assertEqual(select_median_run.operation_rps(summary), 321.5) + self.assertEqual( + select_median_run.operation_rps({"http": {"rps": 123.25}}), + 123.25, + ) + + def test_median_summary_is_paired_with_the_same_runs_resources(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + results_dir = Path(temporary_directory) + (results_dir / "k6").mkdir() + rps_by_run = {1: 120.0, 2: 100.0, 3: 110.0} + for run, rps in rps_by_run.items(): + self._write_json( + results_dir / f"k6_summary_run{run}.json", + { + "marker": f"summary-{run}", + "config": { + "mode": "measurement", + "expected_negotiated_protocol_version": "2025-03-26", + "eligibility_contract": select_median_run.ELIGIBILITY_CONTRACT, + }, + "rates": { + "operations": {"count": 4, "per_second": rps} + }, + "check_breakdown": { + check_name: {"passes": 1, "fails": 0} + for check_name in ( + *select_median_run.OPERATION_SHAPE_CHECKS, + *select_median_run.OPERATION_CONTRACT_CHECKS, + ) + }, + "errors": { + "mcp": 0, + "http": 0, + "checks": 0, + "mcp_rate": 0, + "http_rate": 0, + "check_pass_rate": 1, + }, + }, + ) + samples = [ + { + "timestamp": f"run-{run}-start", + "cpu_percent": run * 10, + "mem_usage_bytes": run * 1_000, + "mem_limit_bytes": 2_000_000, + "net_io_rx": run * 10_000, + "net_io_tx": run * 20_000, + }, + { + "timestamp": f"run-{run}-end", + "cpu_percent": run * 20, + "mem_usage_bytes": run * 2_000, + "mem_limit_bytes": 2_000_000, + "net_io_rx": run * 10_000 + 1_234, + "net_io_tx": run * 20_000 + 5_678, + }, + ] + for prefix in ( + "stats", + "redis_stats", + "api_stats", + "k6_stats", + "host_stats", + ): + self._write_json(results_dir / f"{prefix}_run{run}.json", samples) + audit: dict[str, object] = { + "marker": f"{prefix}-audit-{run}", + "status": "complete", + } + if prefix == "host_stats": + audit["host"] = { + "network_interface_policy": ( + select_median_run.HOST_NETWORK_POLICY + ), + "network_interfaces": ["eth0", "lo"], + "active_network_interfaces": ["eth0", "lo"], + "retired_network_interfaces": [], + "ignored_new_network_interfaces": [], + } + self._write_json( + results_dir / f"{prefix}_run{run}.audit.json", + audit, + ) + for prefix in ("preflight", "postflight"): + self._write_json( + results_dir / f"protocol_{prefix}_run{run}.json", + { + "eligibility_contract": select_median_run.ELIGIBILITY_CONTRACT, + "eligibility_valid": True, + "negotiated_protocol_version": "2025-03-26", + "supplemental_validation": { + "required": False, + "valid": False, + }, + }, + ) + self._write_json( + results_dir / f"resource_headroom_run{run}.json", + {"valid": True}, + ) + self._write_json( + results_dir / f"cgroup_run{run}.json", + {"validation": {"ok": True}}, + ) + self._write_json( + results_dir / "k6" / f"cgroup_run{run}.json", + {"validation": {"ok": True}}, + ) + self._write_json( + results_dir / f"container_inspect_run{run}.json", + {"marker": f"container-{run}"}, + ) + self._write_json( + results_dir / f"image_inspect_run{run}.json", + {"marker": f"image-{run}"}, + ) + self._write_json( + results_dir / "k6" / f"container_inspect_run{run}.json", + {"marker": f"k6-container-{run}"}, + ) + self._write_json( + results_dir / "k6" / f"image_inspect_run{run}.json", + {"marker": f"k6-image-{run}"}, + ) + + with mock.patch.object( + sys, + "argv", + ["select_median_run.py", str(results_dir), "3"], + ): + self.assertEqual(select_median_run.main(), 0) + + self.assertEqual( + self._read_json(results_dir / "k6_summary.json")["marker"], + "summary-3", + ) + self.assertEqual( + self._read_json(results_dir / "stats.json")[0]["timestamp"], + "run-3-start", + ) + self.assertEqual( + self._read_json(results_dir / "container_inspect.json")["marker"], + "container-3", + ) + self.assertEqual( + self._read_json(results_dir / "image_inspect.json")["marker"], + "image-3", + ) + + resource_summary = self._read_json(results_dir / "resource_summary.json") + self.assertEqual(resource_summary["selected_run"], 3) + self.assertEqual( + resource_summary["server"]["network_bytes_during_collection"], + {"rx": 1_234, "tx": 5_678}, + ) + self.assertEqual(resource_summary["server"]["cpu_percent"]["mean"], 45.0) + self.assertEqual( + resource_summary["host"]["network_observation"], + { + "policy": select_median_run.HOST_NETWORK_POLICY, + "scope": "baseline_interface_cohort", + "coverage": "complete", + "observed_bytes_during_collection": { + "rx": 1_234, + "tx": 5_678, + }, + "initial_interfaces": ["eth0", "lo"], + "active_interfaces_at_end": ["eth0", "lo"], + "retired_interfaces": [], + "ignored_new_interfaces": [], + }, + ) + + multi_run = self._read_json(results_dir / "k6_multi_run_stats.json") + self.assertEqual(multi_run["median_run"], 3) + self.assertEqual(multi_run["median_rps"], 110.0) + self.assertEqual( + multi_run["negotiated_protocol_version"], "2025-03-26" + ) + + def test_measurement_rejects_operations_not_reaching_contract_checks(self) -> None: + check_breakdown = { + check_name: {"passes": 1, "fails": 0} + for check_name in ( + *select_median_run.OPERATION_SHAPE_CHECKS, + *select_median_run.OPERATION_CONTRACT_CHECKS, + ) + } + summary = { + "config": { + "mode": "measurement", + "eligibility_contract": select_median_run.ELIGIBILITY_CONTRACT, + }, + "rates": {"operations": {"count": 5, "per_second": 1.0}}, + "check_breakdown": check_breakdown, + "errors": { + "mcp": 0, + "http": 0, + "checks": 0, + "mcp_rate": 0, + "http_rate": 0, + "check_pass_rate": 1, + }, + } + + with self.assertRaisesRegex(ValueError, "only 4 reached the required checks"): + select_median_run.validate_measurement_summary(summary, Path("summary.json")) + + @staticmethod + def _write_json(path: Path, value: object) -> None: + path.write_text(json.dumps(value), encoding="utf-8") + + @staticmethod + def _read_json(path: Path) -> object: + return json.loads(path.read_text(encoding="utf-8")) + + +class CollectorParserTests(unittest.TestCase): + def test_docker_stats_units_and_fields_are_parsed(self) -> None: + self.assertEqual(collect_stats.parse_size_to_bytes("1.5 GiB"), 1_610_612_736) + self.assertEqual(collect_stats.parse_size_to_bytes("2 MB"), 2_000_000) + self.assertEqual( + collect_stats.parse_size_to_bytes("1E+03 MB"), 1_000_000_000 + ) + self.assertEqual( + collect_stats.parse_size_to_bytes("2.5e-1 GiB"), 268_435_456 + ) + self.assertEqual(collect_stats.parse_size_to_bytes("--"), 0) + self.assertEqual(collect_stats.parse_cpu_percent("12.75%"), 12.75) + self.assertEqual( + collect_stats.parse_mem_usage("512 MiB / 2 GiB"), + (536_870_912, 2_147_483_648), + ) + self.assertEqual( + collect_stats.parse_net_io("1.25 MB / 640 kB"), + (1_250_000, 640_000), + ) + self.assertEqual( + collect_stats.parse_net_io("221MB / 1e+03MB"), + (221_000_000, 1_000_000_000), + ) + + def test_malformed_docker_stats_values_are_rejected(self) -> None: + with self.assertRaises(ValueError): + collect_stats.parse_size_to_bytes("forty-two") + with self.assertRaises(ValueError): + collect_stats.parse_size_to_bytes("1eMB") + with self.assertRaises(ValueError): + collect_stats.parse_size_to_bytes("-1 MB") + with self.assertRaises(ValueError): + collect_stats.parse_size_to_bytes("1e3XB") + with self.assertRaises(ValueError): + collect_stats.parse_mem_usage("1 GiB") + with self.assertRaises(ValueError): + collect_stats.parse_net_io("1 MB / 2 MB / 3 MB") + + def test_collector_keeps_list_output_and_writes_audit_sidecar(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + output = Path(temporary_directory) / "stats_run1.json" + call_count = 0 + + class FakeContainerStream: + def __init__(self, _target: str) -> None: + pass + + def collect(self) -> dict[str, int | float | str]: + nonlocal call_count + call_count += 1 + if call_count == 3: + collect_stats.handle_signal(signal.SIGTERM, None) + return { + "timestamp": f"2026-07-19T00:00:0{call_count}+00:00", + "cpu_percent": 10.0, + "mem_usage_bytes": 100, + "mem_limit_bytes": 1_000, + "net_io_rx": call_count, + "net_io_tx": call_count, + } + + def close(self) -> None: + pass + + with ( + mock.patch.object( + sys, + "argv", + ["collect_stats.py", "container", str(output), "1"], + ), + mock.patch.object(collect_stats, "ContainerStatsStream", FakeContainerStream), + mock.patch.object(collect_stats.signal, "signal"), + ): + self.assertEqual(collect_stats.main(), 0) + + samples = json.loads(output.read_text(encoding="utf-8")) + audit = json.loads( + collect_stats.audit_path_for(output).read_text(encoding="utf-8") + ) + self.assertIsInstance(samples, list) + self.assertEqual(len(samples), 3) + self.assertEqual(audit["status"], "complete") + self.assertEqual(audit["attempt_count"], 3) + self.assertEqual(audit["sample_count"], 3) + self.assertEqual(audit["failure_count"], 0) + self.assertEqual(audit["termination"]["signal"], signal.SIGTERM) + self.assertEqual(audit["source"], "persistent docker stats stream") + + def test_host_main_serializes_network_scope_and_churn(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + output = Path(temporary_directory) / "host_stats_run1.json" + + class FakeHostCollector: + def collect(self) -> dict[str, int | float | str]: + collect_stats.handle_signal(signal.SIGTERM, None) + return { + "timestamp": "2026-07-19T00:00:01+00:00", + "cpu_percent": 10.0, + "mem_usage_bytes": 100, + "mem_limit_bytes": 1_000, + "net_io_rx": 10, + "net_io_tx": 20, + } + + def audit_metadata(self) -> dict[str, object]: + return { + "cpu_affinity": [0], + "cpu_capacity_cores": 1, + "network_interfaces": ["eth0", "veth-old"], + "active_network_interfaces": ["eth0"], + "retired_network_interfaces": ["veth-old"], + "ignored_new_network_interfaces": ["veth-new"], + "network_interface_policy": collect_stats.HOST_NETWORK_POLICY, + "network_interface_accounting": "test accounting", + } + + with ( + mock.patch.object( + sys, + "argv", + ["collect_stats.py", "@host", str(output), "1"], + ), + mock.patch.object(collect_stats, "HostStatsCollector", FakeHostCollector), + mock.patch.object(collect_stats, "wait_until"), + mock.patch.object(collect_stats.signal, "signal"), + ): + self.assertEqual(collect_stats.main(), 0) + + audit = json.loads( + collect_stats.audit_path_for(output).read_text(encoding="utf-8") + ) + self.assertEqual(audit["status"], "complete") + self.assertEqual(audit["failure_count"], 0) + self.assertEqual( + audit["host"]["network_interface_policy"], + collect_stats.HOST_NETWORK_POLICY, + ) + self.assertEqual( + audit["host"]["retired_network_interfaces"], ["veth-old"] + ) + self.assertEqual( + audit["host"]["ignored_new_network_interfaces"], ["veth-new"] + ) + + def test_docker_stats_screen_controls_are_not_part_of_the_sample(self) -> None: + line = "\x1b[J\x1b[H12.5%|512 MiB / 2 GiB|221MB / 1e+03MB\x1b[K" + cleaned = collect_stats.ANSI_CONTROL.sub("", line).strip() + sample = collect_stats.parse_container_stats_line(cleaned) + + self.assertEqual(sample["cpu_percent"], 12.5) + self.assertEqual(sample["mem_usage_bytes"], 536_870_912) + self.assertEqual(sample["net_io_rx"], 221_000_000) + self.assertEqual(sample["net_io_tx"], 1_000_000_000) + + def test_host_collector_uses_affinity_cpu_memory_and_network_counters(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + proc_root = Path(temporary_directory) + (proc_root / "net").mkdir() + stat = proc_root / "stat" + stat.write_text( + "cpu 0 0 0 0 0 0 0 0\n" + "cpu0 100 0 50 850 0 0 0 0\n" + "cpu1 200 0 50 750 0 0 0 0\n", + encoding="utf-8", + ) + (proc_root / "meminfo").write_text( + "MemTotal: 1000 kB\nMemAvailable: 400 kB\n", + encoding="utf-8", + ) + (proc_root / "net" / "dev").write_text( + "Inter-| Receive | Transmit\n" + " lo: 10 0 0 0 0 0 0 0 20 0 0 0 0 0 0 0\n" + " eth0: 30 0 0 0 0 0 0 0 40 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + collector = collect_stats.HostStatsCollector( + proc_root, frozenset({0, 1}) + ) + stat.write_text( + "cpu 0 0 0 0 0 0 0 0\n" + "cpu0 125 0 65 910 0 0 0 0\n" + "cpu1 215 0 55 830 0 0 0 0\n", + encoding="utf-8", + ) + + sample = collector.collect() + + self.assertAlmostEqual(float(sample["cpu_percent"]), 60.0) + self.assertEqual(sample["mem_usage_bytes"], 600 * 1024) + self.assertEqual(sample["mem_limit_bytes"], 1000 * 1024) + self.assertEqual(sample["net_io_rx"], 40) + self.assertEqual(sample["net_io_tx"], 60) + + def test_host_collector_retires_disappeared_and_ignores_new_interfaces( + self, + ) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + proc_root = Path(temporary_directory) + (proc_root / "net").mkdir() + stat = proc_root / "stat" + stat.write_text("cpu0 100 0 50 850 0 0 0 0\n", encoding="utf-8") + (proc_root / "meminfo").write_text( + "MemTotal: 1000 kB\nMemAvailable: 400 kB\n", + encoding="utf-8", + ) + network = proc_root / "net" / "dev" + network.write_text( + "Inter-| Receive | Transmit\n" + " lo: 10 0 0 0 0 0 0 0 20 0 0 0 0 0 0 0\n" + " eth0: 30 0 0 0 0 0 0 0 40 0 0 0 0 0 0 0\n" + " veth-old: 500 0 0 0 0 0 0 0 600 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + collector = collect_stats.HostStatsCollector( + proc_root, frozenset({0}) + ) + stat.write_text("cpu0 110 0 55 885 0 0 0 0\n", encoding="utf-8") + network.write_text( + "Inter-| Receive | Transmit\n" + " lo: 12 0 0 0 0 0 0 0 23 0 0 0 0 0 0 0\n" + " eth0: 35 0 0 0 0 0 0 0 45 0 0 0 0 0 0 0\n" + " veth-new: 1000 0 0 0 0 0 0 0 2000 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + + sample = collector.collect() + + self.assertEqual(sample["net_io_rx"], 547) + self.assertEqual(sample["net_io_tx"], 668) + self.assertEqual( + collector.network_interfaces, + frozenset({"lo", "eth0", "veth-old"}), + ) + self.assertEqual(collector.active_network_interfaces, {"lo", "eth0"}) + self.assertEqual(collector.retired_network_interfaces, {"veth-old"}) + + # Reusing the retired name with lower counters must not create a + # false reset or contaminate the original baseline cohort. + stat.write_text("cpu0 120 0 60 920 0 0 0 0\n", encoding="utf-8") + network.write_text( + "Inter-| Receive | Transmit\n" + " lo: 14 0 0 0 0 0 0 0 26 0 0 0 0 0 0 0\n" + " eth0: 38 0 0 0 0 0 0 0 48 0 0 0 0 0 0 0\n" + " veth-old: 1 0 0 0 0 0 0 0 2 0 0 0 0 0 0 0\n" + " veth-new: 2000 0 0 0 0 0 0 0 3000 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + + sample = collector.collect() + + self.assertEqual(sample["net_io_rx"], 552) + self.assertEqual(sample["net_io_tx"], 674) + self.assertEqual(collector.active_network_interfaces, {"lo", "eth0"}) + self.assertEqual(collector.retired_network_interfaces, {"veth-old"}) + self.assertEqual( + collector.ignored_new_network_interfaces, + {"veth-new"}, + ) + self.assertEqual( + collector.audit_metadata()["retired_network_interfaces"], + ["veth-old"], + ) + self.assertEqual( + collector.audit_metadata()["ignored_new_network_interfaces"], + ["veth-new"], + ) + self.assertEqual( + collector.audit_metadata()["network_interface_policy"], + collect_stats.HOST_NETWORK_POLICY, + ) + self.assertIn( + "freeze last counters", + collector.audit_metadata()["network_interface_accounting"], + ) + + def test_host_collector_rejects_resets_on_selected_interfaces(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + proc_root = Path(temporary_directory) + (proc_root / "net").mkdir() + stat = proc_root / "stat" + stat.write_text("cpu0 100 0 50 850 0 0 0 0\n", encoding="utf-8") + (proc_root / "meminfo").write_text( + "MemTotal: 1000 kB\nMemAvailable: 400 kB\n", + encoding="utf-8", + ) + network = proc_root / "net" / "dev" + network.write_text( + "Inter-| Receive | Transmit\n" + " lo: 10 0 0 0 0 0 0 0 20 0 0 0 0 0 0 0\n" + " eth0: 30 0 0 0 0 0 0 0 40 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + collector = collect_stats.HostStatsCollector( + proc_root, frozenset({0}) + ) + stat.write_text("cpu0 120 0 60 920 0 0 0 0\n", encoding="utf-8") + network.write_text( + "Inter-| Receive | Transmit\n" + " lo: 9 0 0 0 0 0 0 0 24 0 0 0 0 0 0 0\n" + " eth0: 36 0 0 0 0 0 0 0 46 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + with self.assertRaisesRegex(RuntimeError, "counters reset"): + collector.collect() + + def test_host_collector_fails_if_every_baseline_interface_retires(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + proc_root = Path(temporary_directory) + (proc_root / "net").mkdir() + stat = proc_root / "stat" + stat.write_text("cpu0 100 0 50 850 0 0 0 0\n", encoding="utf-8") + (proc_root / "meminfo").write_text( + "MemTotal: 1000 kB\nMemAvailable: 400 kB\n", + encoding="utf-8", + ) + network = proc_root / "net" / "dev" + network.write_text( + "Inter-| Receive | Transmit\n" + " old0: 10 0 0 0 0 0 0 0 20 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + collector = collect_stats.HostStatsCollector( + proc_root, frozenset({0}) + ) + stat.write_text("cpu0 110 0 55 885 0 0 0 0\n", encoding="utf-8") + network.write_text( + "Inter-| Receive | Transmit\n" + " new0: 30 0 0 0 0 0 0 0 40 0 0 0 0 0 0 0\n", + encoding="utf-8", + ) + + with self.assertRaisesRegex( + RuntimeError, + "all baseline host network interfaces disappeared", + ): + collector.collect() + + +class CollectorFailureSummaryTests(unittest.TestCase): + def test_failed_audit_is_used_instead_of_a_later_empty_log(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + results_dir = Path(temporary_directory) + api_audit = results_dir / "api_stats_run1.audit.json" + host_audit = results_dir / "host_stats_run1.audit.json" + api_audit.write_text( + json.dumps( + { + "status": "failed", + "target": "mcp-api-service", + "failures": [ + {"error": "Unable to parse size value: '1e+03MB'"} + ], + } + ), + encoding="utf-8", + ) + host_audit.write_text( + json.dumps({"status": "complete", "target": "@host"}), + encoding="utf-8", + ) + + summary = summarize_collector_failures.summarize_failures( + [api_audit, host_audit] + ) + + self.assertEqual( + summary, + "mcp-api-service: Unable to parse size value: '1e+03MB'", + ) + + def test_missing_audit_is_reported(self) -> None: + missing = Path("missing_stats_run1.audit.json") + + summary = summarize_collector_failures.summarize_failures([missing]) + + self.assertIn("missing_stats_run1.audit.json: unavailable audit", summary) + + +class ResourceHeadroomTests(unittest.TestCase): + MEMORY_LIMIT = 1_000 + + def samples( + self, + *, + count: int = 5, + step_seconds: float = 1.0, + cpu_percent: float = 25.0, + memory_limit: int = MEMORY_LIMIT, + ) -> list[dict[str, int | float | str]]: + start = datetime(2026, 7, 19, tzinfo=timezone.utc) + return [ + { + "timestamp": (start + timedelta(seconds=index * step_seconds)).isoformat(), + "cpu_percent": cpu_percent, + "mem_usage_bytes": 250, + "mem_limit_bytes": memory_limit, + "net_io_rx": index, + "net_io_tx": index, + } + for index in range(count) + ] + + @staticmethod + def audit(sample_count: int = 5, **overrides: object) -> dict[str, object]: + value: dict[str, object] = { + "schema_version": 1, + "status": "complete", + "interval_seconds": 1.0, + "elapsed_seconds": 5.0, + "termination": {"reason": "signal", "signal": signal.SIGTERM}, + "attempt_count": sample_count, + "sample_count": sample_count, + "failure_count": 0, + "failures": [], + "fatal_error": None, + } + value.update(overrides) + return value + + def evaluate( + self, + samples: list[dict[str, int | float | str]], + audit: dict[str, object], + *, + enforce_headroom: bool = True, + ) -> dict[str, object]: + return validate_resource_headroom.evaluate( + "resource", + samples, + audit, + cpu_limit=1.0, + memory_limit=self.MEMORY_LIMIT, + threshold=0.90, + expected_duration=5.0, + enforce_headroom=enforce_headroom, + ) + + def test_failed_collection_sparse_coverage_and_limit_mismatch_are_rejected(self) -> None: + failed = self.evaluate( + self.samples(), + self.audit( + attempt_count=6, + failure_count=1, + status="failed", + ), + ) + sparse = self.evaluate( + self.samples(step_seconds=2.0), + self.audit(), + ) + wrong_limit = self.evaluate( + self.samples(memory_limit=999), + self.audit(), + ) + + self.assertFalse(failed["valid"]) + self.assertIn("failed attempts", str(failed["reason"])) + self.assertFalse(sparse["valid"]) + self.assertIn("maximum sample gap", str(sparse["reason"])) + self.assertFalse(wrong_limit["valid"]) + self.assertIn("mem_limit_bytes=999", str(wrong_limit["reason"])) + + def test_saturated_observed_target_is_valid_but_shared_resource_is_not(self) -> None: + samples = self.samples(cpu_percent=95.0) + audit = self.audit() + + shared = self.evaluate(samples, audit, enforce_headroom=True) + observed = self.evaluate(samples, audit, enforce_headroom=False) + + self.assertFalse(shared["valid"]) + self.assertTrue(shared["headroom_exceeded"]) + self.assertTrue(observed["valid"]) + self.assertTrue(observed["headroom_exceeded"]) + self.assertFalse(observed["headroom_enforced"]) + + def test_duration_parser_accepts_seconds_and_k6_style_suffixes(self) -> None: + self.assertEqual(validate_resource_headroom.parse_duration("300"), 300.0) + self.assertEqual(validate_resource_headroom.parse_duration("300s"), 300.0) + self.assertEqual(validate_resource_headroom.parse_duration("5m"), 300.0) + + def test_decreasing_network_counters_are_rejected(self) -> None: + samples = self.samples() + samples[3]["net_io_rx"] = 1 + + result = self.evaluate(samples, self.audit()) + + self.assertFalse(result["valid"]) + self.assertIn("network counters decreased", str(result["reason"])) + + def test_resource_summary_does_not_mask_decreasing_network_counters(self) -> None: + samples = self.samples() + samples[3]["net_io_tx"] = 1 + + with self.assertRaisesRegex(ValueError, "not monotonic"): + select_median_run.resource_summary(samples, selected_run=1) + + def test_host_resource_summary_labels_partial_cohort_coverage(self) -> None: + summary = select_median_run.resource_summary( + self.samples(), + selected_run=1, + host_audit={ + "status": "complete", + "host": { + "network_interface_policy": ( + select_median_run.HOST_NETWORK_POLICY + ), + "network_interfaces": ["eth0", "veth-old"], + "active_network_interfaces": ["eth0"], + "retired_network_interfaces": ["veth-old"], + "ignored_new_network_interfaces": ["veth-new"], + }, + }, + ) + + self.assertNotIn("network_bytes_during_collection", summary) + self.assertEqual( + summary["network_observation"], + { + "policy": select_median_run.HOST_NETWORK_POLICY, + "scope": "baseline_interface_cohort", + "coverage": "partial", + "observed_bytes_during_collection": {"rx": 4, "tx": 4}, + "initial_interfaces": ["eth0", "veth-old"], + "active_interfaces_at_end": ["eth0"], + "retired_interfaces": ["veth-old"], + "ignored_new_interfaces": ["veth-new"], + }, + ) + + +class ProvenanceTests(unittest.TestCase): + def test_tree_digest_tracks_sources_but_not_result_artifacts(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + project_dir = Path(temporary_directory) + (project_dir / "include").mkdir() + (project_dir / "benchmark" / "alternatives").mkdir(parents=True) + (project_dir / "benchmark" / "results").mkdir(parents=True) + (project_dir / "CMakeLists.txt").write_text("project(test)\n", encoding="utf-8") + header = project_dir / "include" / "api.hpp" + header.write_text("void api();\n", encoding="utf-8") + adapter = project_dir / "benchmark" / "alternatives" / "adapter.cpp" + adapter.write_text("int adapter();\n", encoding="utf-8") + result = project_dir / "benchmark" / "results" / "summary.json" + result.write_text("{}\n", encoding="utf-8") + + initial_digest, initial_paths = capture_environment.tree_digest(project_dir) + self.assertEqual(len(initial_paths), 3) + + result.write_text('{"changed": true}\n', encoding="utf-8") + unchanged_digest, unchanged_paths = capture_environment.tree_digest(project_dir) + self.assertEqual((unchanged_digest, unchanged_paths), (initial_digest, initial_paths)) + + adapter.write_text("int changed_adapter();\n", encoding="utf-8") + changed_digest, changed_paths = capture_environment.tree_digest(project_dir) + self.assertEqual(changed_paths, initial_paths) + self.assertNotEqual(changed_digest, initial_digest) + + def test_source_snapshot_is_reconstructible_and_change_check_fails_closed(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + project_dir = Path(temporary_directory) + (project_dir / "include").mkdir() + (project_dir / "benchmark" / "results").mkdir(parents=True) + (project_dir / "CMakeLists.txt").write_text("project(test)\n", encoding="utf-8") + source = project_dir / "include" / "api.hpp" + source.write_text("void api();\n", encoding="utf-8") + digest, paths = capture_environment.tree_digest(project_dir) + snapshot = project_dir / "benchmark" / "results" / "source_snapshot.tar" + snapshot_digest = capture_environment.create_source_snapshot( + project_dir, paths, snapshot + ) + self.assertEqual(len(snapshot_digest), 64) + environment = project_dir / "benchmark" / "results" / "environment.json" + environment.write_text( + json.dumps({"source": {"tree_sha256": digest}}), encoding="utf-8" + ) + self.assertEqual( + capture_environment.verify_source_digest(environment, project_dir), 0 + ) + source.write_text("void changed();\n", encoding="utf-8") + self.assertEqual( + capture_environment.verify_source_digest(environment, project_dir), 1 + ) + + def test_commented_alternative_manifest_header_is_accepted(self) -> None: + with tempfile.TemporaryDirectory() as temporary_directory: + root = Path(temporary_directory) + manifest = root / "sources.tsv" + manifest.write_text( + "# id\trepository\tcommit\tbenchmark_status\n" + "sdk\thttps://example.test/sdk.git\tabc123\tsupported\n", + encoding="utf-8", + ) + with mock.patch.object( + capture_environment, + "git_repository", + return_value={"origin_url": "https://example.test/sdk.git"}, + ): + repositories = capture_environment.alternative_repositories( + manifest, root + ) + self.assertEqual(repositories[0]["id"], "sdk") + self.assertTrue(repositories[0]["origin_matches_manifest"]) + + +class ContainerContractTests(unittest.TestCase): + def test_fractional_cpu_limit_is_validated_in_docker_and_cgroup(self) -> None: + inspect = { + "HostConfig": { + "NanoCpus": 500_000_000, + "Memory": 536_870_912, + "CpusetCpus": "2-3", + } + } + selected = lambda value: {"selected": {"value": value}, "attempts": []} + cgroup = { + "version": 2, + "cpu": selected("50000 100000"), + "memory": selected("536870912"), + "cpuset_effective": selected("2,3"), + } + result = capture_container.validate_resources( + inspect, + cgroup, + expected_cpus=0.5, + expected_memory_bytes=536_870_912, + expected_cpuset="2-3", + ) + self.assertTrue(result["ok"]) + + def test_relative_executable_is_resolved_from_container_workdir(self) -> None: + commands: list[list[str]] = [] + + def run(command: list[str]) -> mock.Mock: + commands.append(command) + if command[:3] == ["docker", "exec", "api"]: + return mock.Mock(returncode=0, stdout="/app/./api-service\n", stderr="") + if command[:3] == ["docker", "cp", "-L"]: + Path(command[-1]).write_bytes(b"server executable") + return mock.Mock(returncode=0, stdout="", stderr="") + self.fail(f"unexpected command: {command!r}") + + with mock.patch.object(capture_container, "run_command", side_effect=run): + provenance = capture_container.executable_provenance("api", "./api-service") + + self.assertEqual(provenance["resolved_path"], "/app/./api-service") + self.assertIsNotNone(provenance["sha256"]) + self.assertIn("api:/app/./api-service", commands[1]) + + +class ProtocolCorrectnessTests(unittest.TestCase): + class _FakeResponse: + def __init__( + self, + body: bytes = b"", + headers: dict[str, str] | None = None, + status: int = 200, + ) -> None: + self._body = body + self.headers = headers if headers is not None else { + "Content-Type": "application/json" + } + self.status = status + + def read(self) -> bytes: + return self._body + + def getcode(self) -> int: + return self.status + + def __enter__(self) -> "ProtocolCorrectnessTests._FakeResponse": + return self + + def __exit__(self, *_: object) -> None: + return None + + def test_inherited_contract_accepts_python_schema_and_rust_counter(self) -> None: + python_checkout_tool = { + "name": "checkout", + "inputSchema": { + "type": "object", + "properties": { + "user_id": {"type": "string"}, + "items": {"type": "array", "items": {}}, + }, + }, + } + rust_checkout = { + "user_id": "user-00001", + "status": "confirmed", + "total": 24.63, + "items_count": 2, + "rate_limit_count": 295, + } + + self.assertIsInstance(python_checkout_tool["inputSchema"], dict) + self.assertFalse(verify_server.validate_tool_schema(python_checkout_tool)) + self.assertTrue( + verify_server.valid_upstream_checkout(rust_checkout, "user-00001") + ) + self.assertFalse(verify_server.valid_exact_checkout(rust_checkout, "rust")) + + def test_checkout_side_effect_gate_checks_all_three_redis_mutations(self) -> None: + before = { + "rate_limit_count": 0, + "history_length": 20, + "product_42_popularity": 294.0, + } + after = { + "rate_limit_count": 1, + "history_length": 21, + "product_42_popularity": 295.0, + } + + self.assertTrue(verify_server.valid_workload_side_effects(before, after)) + after["history_length"] = 20 + self.assertFalse(verify_server.valid_workload_side_effects(before, after)) + + def test_schema_contract_allows_metadata_and_closed_objects(self) -> None: + tool = { + "name": "checkout", + "inputSchema": { + "type": "object", + "additionalProperties": False, + "properties": { + "user_id": {"type": "string"}, + "items": { + "type": "array", + "items": { + "type": "object", + "additionalProperties": False, + "required": ["quantity", "product_id"], + "properties": { + "quantity": {"type": "integer"}, + "product_id": {"type": "integer"}, + }, + }, + }, + }, + }, + } + + self.assertTrue(verify_server.validate_tool_schema(tool)) + + def test_schema_contract_rejects_unmodeled_narrowing_assertions(self) -> None: + tool = { + "name": "checkout", + "inputSchema": { + "type": "object", + "properties": { + "user_id": {"type": "string"}, + "items": { + "type": "array", + "items": { + "type": "object", + "properties": { + "product_id": { + "type": "integer", + "maximum": 0, + }, + "quantity": {"type": "integer"}, + }, + "required": ["product_id", "quantity"], + }, + }, + }, + }, + } + + self.assertFalse(verify_server.validate_tool_schema(tool)) + + def test_schema_contract_accepts_compatible_rust_style_schema(self) -> None: + tool = { + "name": "checkout", + "inputSchema": { + "$schema": "https://json-schema.org/draft/2020-12/schema", + "type": "object", + "$defs": { + "CheckoutItem": { + "type": "object", + "properties": { + "product_id": { + "type": "integer", + "format": "uint32", + "minimum": 0, + }, + "quantity": { + "type": "integer", + "format": "uint32", + "minimum": 0, + }, + }, + "required": ["product_id", "quantity"], + } + }, + "properties": { + "user_id": {"type": "string"}, + "items": { + "type": "array", + "items": {"$ref": "#/$defs/CheckoutItem"}, + }, + }, + }, + } + + self.assertTrue(verify_server.validate_tool_schema(tool)) + + def test_schema_contract_rejects_bound_that_excludes_fixture(self) -> None: + tool = { + "name": "checkout", + "inputSchema": { + "type": "object", + "properties": { + "user_id": {"type": "string"}, + "items": { + "type": "array", + "items": { + "type": "object", + "properties": { + "product_id": {"type": "integer", "minimum": 43}, + "quantity": {"type": "integer"}, + }, + "required": ["product_id", "quantity"], + }, + }, + }, + }, + } + + self.assertFalse(verify_server.validate_tool_schema(tool)) + + def test_schema_contract_rejects_unknown_assertion(self) -> None: + tool = { + "name": "get_user_cart", + "inputSchema": { + "type": "object", + "properties": { + "user_id": { + "type": "string", + "x-requires-canonical-user": True, + } + }, + }, + } + + self.assertFalse(verify_server.validate_tool_schema(tool)) + + def test_schema_contract_rejects_unresolved_or_external_references(self) -> None: + for reference in ("#/$defs/Missing", "https://example.com/item.json"): + with self.subTest(reference=reference): + tool = { + "name": "checkout", + "inputSchema": { + "type": "object", + "properties": { + "user_id": {"type": "string"}, + "items": { + "type": "array", + "items": {"$ref": reference}, + }, + }, + }, + } + self.assertFalse(verify_server.validate_tool_schema(tool)) + + def test_response_parser_accepts_json_and_sse(self) -> None: + self.assertEqual( + verify_server.parse_response(b'{"result": 1}', "application/json"), + {"result": 1}, + ) + self.assertEqual( + verify_server.parse_response( + b"event: message\ndata: {\"result\": 2}\n\n", + "text/event-stream", + ), + {"result": 2}, + ) + self.assertIsNone(verify_server.parse_response(b" \n", "application/json")) + with self.assertRaises(RuntimeError): + verify_server.parse_response(b"not an MCP response", "text/event-stream") + with self.assertRaises(json.JSONDecodeError): + verify_server.parse_response( + b"data: {\"result\": 2}\n\n", "application/json" + ) + with self.assertRaises(RuntimeError): + verify_server.parse_response(b"[{\"result\": 1}]", "application/json") + + def test_negotiated_protocol_and_session_headers_propagate(self) -> None: + requests: list[object] = [] + responses = iter( + [ + self._FakeResponse( + b'{"jsonrpc":"2.0","id":1,"result":' + b'{"protocolVersion":"2025-03-26",' + b'"capabilities":{"tools":{}},' + b'"serverInfo":{"name":"benchmark","version":"1.0"}}}', + { + "Mcp-Session-Id": "session-123", + "Content-Type": "application/json", + }, + ), + self._FakeResponse(headers={}, status=202), + self._FakeResponse(b'{"jsonrpc":"2.0","id":2,"result":{"tools":[]}}'), + self._FakeResponse(headers={}, status=204), + ] + ) + + def fake_urlopen(request: object, timeout: int) -> object: + requests.append(request) + return next(responses) + + with mock.patch.object( + verify_server.urllib.request, + "urlopen", + side_effect=fake_urlopen, + ): + session = verify_server.McpSession("http://benchmark.test/mcp") + session.post( + {"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}} + ) + session.close() + + self.assertEqual(len(requests), 4) + request_headers = [ + {key.lower(): value for key, value in request.header_items()} + for request in requests + ] + self.assertNotIn("mcp-protocol-version", request_headers[0]) + for headers in request_headers[1:]: + self.assertEqual(headers["mcp-protocol-version"], "2025-03-26") + self.assertEqual(headers["mcp-session-id"], "session-123") + self.assertEqual(requests[0].get_method(), "POST") + self.assertEqual(requests[-1].get_method(), "DELETE") + + initialize_payload = json.loads(requests[0].data.decode("utf-8")) + self.assertEqual( + initialize_payload["params"]["protocolVersion"], + verify_server.PROTOCOL_VERSION, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/benchmark/validate_resource_headroom.py b/benchmark/validate_resource_headroom.py new file mode 100644 index 0000000..73a81cd --- /dev/null +++ b/benchmark/validate_resource_headroom.py @@ -0,0 +1,440 @@ +#!/usr/bin/env python3 +"""Validate resource-sampling integrity and shared-resource headroom.""" + +from __future__ import annotations + +import argparse +import json +import math +import re +from datetime import datetime +from pathlib import Path +from typing import Any + + +SAMPLE_FIELDS = { + "timestamp", + "cpu_percent", + "mem_usage_bytes", + "mem_limit_bytes", + "net_io_rx", + "net_io_tx", +} + + +def percentile(values: list[float], fraction: float) -> float: + """Return a nearest-rank percentile without interpolating observations.""" + if not values: + raise ValueError("at least one sample is required") + rank = max(1, math.ceil(len(values) * fraction)) + return sorted(values)[rank - 1] + + +def parse_duration(value: str) -> float: + match = re.fullmatch(r"\s*([0-9]*\.?[0-9]+)\s*(ms|s|m|h)?\s*", value) + if not match: + raise argparse.ArgumentTypeError( + "duration must be seconds or use an ms, s, m, or h suffix" + ) + units = {None: 1.0, "ms": 0.001, "s": 1.0, "m": 60.0, "h": 3600.0} + seconds = float(match.group(1)) * units[match.group(2)] + if not math.isfinite(seconds) or seconds <= 0: + raise argparse.ArgumentTypeError("duration must be greater than zero") + return seconds + + +def audit_path_for(sample_path: Path) -> Path: + return sample_path.with_name(f"{sample_path.stem}.audit.json") + + +def parse_timestamp(value: Any) -> datetime: + if not isinstance(value, str): + raise ValueError("timestamp is not a string") + try: + timestamp = datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError as error: + raise ValueError(f"invalid ISO-8601 timestamp {value!r}") from error + if timestamp.tzinfo is None or timestamp.utcoffset() is None: + raise ValueError(f"timestamp lacks a UTC offset: {value!r}") + return timestamp + + +def exact_nonnegative_integer(value: Any, description: str) -> int: + if ( + isinstance(value, bool) + or not isinstance(value, (int, float)) + or not math.isfinite(float(value)) + or float(value) != int(value) + ): + raise ValueError(f"{description} must be an exact integer") + parsed = int(value) + if parsed < 0: + raise ValueError(f"{description} must be non-negative") + return parsed + + +def sampling_integrity( + samples: list[dict[str, Any]], + audit: dict[str, Any], + memory_limit: int, + expected_duration: float, +) -> dict[str, Any]: + """Check audit status, schema, declared limits, and temporal coverage.""" + reasons: list[str] = [] + try: + interval = float(audit["interval_seconds"]) + except (KeyError, TypeError, ValueError): + interval = math.nan + reasons.append("audit metadata has no valid interval_seconds") + if not math.isfinite(interval) or interval <= 0: + if not reasons: + reasons.append("audit interval_seconds must be finite and positive") + interval = math.nan + + if audit.get("schema_version") != 1: + reasons.append("unsupported or missing audit schema_version") + if audit.get("status") != "complete": + reasons.append(f"collector status is {audit.get('status')!r}, not 'complete'") + failure_count = audit.get("failure_count") + if isinstance(failure_count, bool) or failure_count != 0: + reasons.append(f"collector recorded {failure_count!r} failed attempts") + recorded_failures = audit.get("failures") + if not isinstance(recorded_failures, list): + reasons.append("collector audit has no failures list") + elif isinstance(failure_count, int) and not isinstance(failure_count, bool) \ + and len(recorded_failures) != failure_count: + reasons.append("collector failure_count does not match its failures list") + if audit.get("fatal_error") is not None: + reasons.append(f"collector recorded fatal error: {audit.get('fatal_error')}") + termination = audit.get("termination") + if not isinstance(termination, dict) or termination.get("reason") != "signal": + reasons.append("collector did not record a clean signal-driven stop") + if audit.get("sample_count") != len(samples): + reasons.append( + "audit sample_count does not match the sample file " + f"({audit.get('sample_count')!r} != {len(samples)})" + ) + if audit.get("attempt_count") != len(samples): + reasons.append( + "collector attempts do not exactly match successful samples " + f"({audit.get('attempt_count')!r} != {len(samples)})" + ) + + timestamps: list[datetime] = [] + cpu: list[float] = [] + memory: list[int] = [] + network: list[tuple[int, int, int]] = [] + for index, sample in enumerate(samples): + if not isinstance(sample, dict): + reasons.append(f"sample {index} is not an object") + continue + missing = SAMPLE_FIELDS - sample.keys() + if missing: + reasons.append(f"sample {index} is missing fields: {sorted(missing)}") + continue + try: + timestamps.append(parse_timestamp(sample["timestamp"])) + if ( + isinstance(sample["cpu_percent"], bool) + or not isinstance(sample["cpu_percent"], (int, float)) + ): + raise ValueError("cpu_percent must be numeric and not boolean") + cpu_value = float(sample["cpu_percent"]) + memory_value = exact_nonnegative_integer( + sample["mem_usage_bytes"], "mem_usage_bytes" + ) + sample_memory_limit = exact_nonnegative_integer( + sample["mem_limit_bytes"], "mem_limit_bytes" + ) + net_rx = exact_nonnegative_integer(sample["net_io_rx"], "net_io_rx") + net_tx = exact_nonnegative_integer(sample["net_io_tx"], "net_io_tx") + if not math.isfinite(cpu_value) or cpu_value < 0: + raise ValueError("cpu_percent must be finite and non-negative") + if memory_value > memory_limit: + raise ValueError("mem_usage_bytes exceeds the declared memory limit") + cpu.append(cpu_value) + memory.append(memory_value) + network.append((index, net_rx, net_tx)) + if ( + sample_memory_limit != memory_limit + ): + reasons.append( + f"sample {index} reports mem_limit_bytes={sample_memory_limit}, " + f"expected {memory_limit}" + ) + except (TypeError, ValueError) as error: + reasons.append(f"sample {index} has invalid values: {error}") + + for earlier, later in zip(network, network[1:]): + earlier_index, earlier_rx, earlier_tx = earlier + later_index, later_rx, later_tx = later + if later_rx < earlier_rx or later_tx < earlier_tx: + reasons.append( + "network counters decreased between samples " + f"{earlier_index} and {later_index}" + ) + + minimum_samples = 3 + if math.isfinite(interval): + minimum_samples = max(3, math.floor(expected_duration / interval) - 1) + if len(samples) < minimum_samples: + reasons.append( + f"only {len(samples)} samples were collected; expected at least " + f"{minimum_samples} for {expected_duration:g}s" + ) + + span_seconds = 0.0 + maximum_gap_seconds = 0.0 + if len(timestamps) >= 2: + gaps = [ + (later - earlier).total_seconds() + for earlier, later in zip(timestamps, timestamps[1:]) + ] + if any(gap <= 0 for gap in gaps): + reasons.append("sample timestamps are not strictly increasing") + span_seconds = (timestamps[-1] - timestamps[0]).total_seconds() + maximum_gap_seconds = max(gaps) + if math.isfinite(interval): + minimum_span = max(0.0, expected_duration - 2.0 * interval) + if span_seconds < minimum_span: + reasons.append( + f"sample timestamps cover only {span_seconds:.3f}s; " + f"expected at least {minimum_span:.3f}s" + ) + maximum_allowed_gap = interval * 1.75 + if maximum_gap_seconds > maximum_allowed_gap: + reasons.append( + f"maximum sample gap is {maximum_gap_seconds:.3f}s; " + f"allowed at most {maximum_allowed_gap:.3f}s" + ) + elif samples: + reasons.append("fewer than two valid timestamps were collected") + + elapsed = audit.get("elapsed_seconds") + try: + elapsed_seconds = float(elapsed) + except (TypeError, ValueError): + elapsed_seconds = math.nan + reasons.append("audit metadata has no valid elapsed_seconds") + if not math.isfinite(elapsed_seconds) or elapsed_seconds < 0: + if not any("elapsed_seconds" in reason for reason in reasons): + reasons.append("audit elapsed_seconds must be finite and non-negative") + elapsed_seconds = math.nan + if math.isfinite(elapsed_seconds) and math.isfinite(interval): + minimum_elapsed = max(0.0, expected_duration - interval) + if elapsed_seconds < minimum_elapsed: + reasons.append( + f"collector ran for only {elapsed_seconds:.3f}s; " + f"expected at least {minimum_elapsed:.3f}s" + ) + + return { + "valid": not reasons, + "reason": "; ".join(reasons) if reasons else None, + "sample_count": len(samples), + "minimum_sample_count": minimum_samples, + "expected_duration_seconds": expected_duration, + "interval_seconds": interval if math.isfinite(interval) else None, + "timestamp_span_seconds": span_seconds, + "maximum_timestamp_gap_seconds": maximum_gap_seconds, + "audit_elapsed_seconds": elapsed_seconds if math.isfinite(elapsed_seconds) else None, + "cpu_values": cpu, + "memory_values": memory, + } + + +def evaluate( + name: str, + samples: list[dict[str, Any]], + audit: dict[str, Any], + cpu_limit: float, + memory_limit: int, + threshold: float, + expected_duration: float, + enforce_headroom: bool, +) -> dict[str, Any]: + integrity = sampling_integrity(samples, audit, memory_limit, expected_duration) + cpu = integrity.pop("cpu_values") + memory = integrity.pop("memory_values") + result: dict[str, Any] = { + "name": name, + "role": "shared" if enforce_headroom else "observed_target", + "headroom_enforced": enforce_headroom, + "collection": integrity, + } + + if cpu and memory: + cpu_capacity = cpu_limit * 100.0 + cpu_p95 = percentile(cpu, 0.95) + memory_p95 = percentile([float(value) for value in memory], 0.95) + cpu_ratio = cpu_p95 / cpu_capacity + memory_ratio = memory_p95 / memory_limit + headroom_reasons = [] + if cpu_ratio >= threshold: + headroom_reasons.append(f"p95 CPU used {cpu_ratio:.1%} of its limit") + if memory_ratio >= threshold: + headroom_reasons.append(f"p95 memory used {memory_ratio:.1%} of its limit") + result.update( + { + "cpu": { + "limit_cores": cpu_limit, + "p95_percent": cpu_p95, + "p95_limit_ratio": cpu_ratio, + }, + "memory": { + "limit_bytes": memory_limit, + "p95_bytes": int(memory_p95), + "p95_limit_ratio": memory_ratio, + }, + "headroom_exceeded": bool(headroom_reasons), + "headroom_reason": "; ".join(headroom_reasons) or None, + } + ) + else: + result["headroom_exceeded"] = None + result["headroom_reason"] = "no valid samples available for p95 calculation" + + headroom_valid = not result["headroom_exceeded"] if enforce_headroom else True + result["valid"] = bool(integrity["valid"] and headroom_valid) + reasons = [integrity["reason"]] if integrity["reason"] else [] + if enforce_headroom and result["headroom_reason"]: + reasons.append(result["headroom_reason"]) + result["reason"] = "; ".join(reasons) or None + return result + + +def parse_resource(value: str) -> tuple[str, Path, float, int]: + try: + name, path, cpu_limit, memory_limit = value.split(":", 3) + parsed_cpu = float(cpu_limit) + parsed_memory = int(memory_limit) + if ( + not name + or not math.isfinite(parsed_cpu) + or parsed_cpu <= 0 + or parsed_memory <= 0 + ): + raise ValueError + return name, Path(path), parsed_cpu, parsed_memory + except (TypeError, ValueError) as error: + raise argparse.ArgumentTypeError( + "resource must be NAME:PATH:CPU_CORES:MEMORY_BYTES with positive limits" + ) from error + + +def load_resource( + name: str, + path: Path, + cpu_limit: float, + memory_limit: int, + threshold: float, + expected_duration: float, + enforce_headroom: bool, +) -> dict[str, Any]: + audit_path = audit_path_for(path) + try: + with path.open(encoding="utf-8") as stream: + samples = json.load(stream) + if not isinstance(samples, list): + raise ValueError("sample JSON must contain a top-level list") + with audit_path.open(encoding="utf-8") as stream: + audit = json.load(stream) + if not isinstance(audit, dict): + raise ValueError("audit JSON must contain a top-level object") + except (OSError, ValueError, json.JSONDecodeError) as error: + return { + "name": name, + "role": "shared" if enforce_headroom else "observed_target", + "headroom_enforced": enforce_headroom, + "valid": False, + "reason": f"unable to load samples and audit metadata: {error}", + "sample_path": str(path), + "audit_path": str(audit_path), + } + + result = evaluate( + name, + samples, + audit, + cpu_limit, + memory_limit, + threshold, + expected_duration, + enforce_headroom, + ) + result["sample_path"] = str(path) + result["audit_path"] = str(audit_path) + return result + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("output", type=Path) + parser.add_argument( + "--resource", + action="append", + type=parse_resource, + default=[], + help="shared NAME:PATH:CPU_CORES:MEMORY_BYTES (p95 headroom enforced)", + ) + parser.add_argument( + "--observed-resource", + action="append", + type=parse_resource, + default=[], + help="target NAME:PATH:CPU_CORES:MEMORY_BYTES (collection integrity only)", + ) + parser.add_argument( + "--expected-duration", + "--expected-duration-seconds", + dest="expected_duration", + type=parse_duration, + required=True, + help="expected collection duration, for example 300, 300s, or 5m", + ) + parser.add_argument("--threshold", type=float, default=0.90) + args = parser.parse_args() + if not 0 < args.threshold < 1: + parser.error("--threshold must be between 0 and 1") + if not args.resource and not args.observed_resource: + parser.error("at least one --resource or --observed-resource is required") + + definitions = [ + (*resource, True) for resource in args.resource + ] + [ + (*resource, False) for resource in args.observed_resource + ] + names = [definition[0] for definition in definitions] + if len(names) != len(set(names)): + parser.error("resource names must be unique") + + results = [ + load_resource( + name, + path, + cpu_limit, + memory_limit, + args.threshold, + args.expected_duration, + enforce_headroom, + ) + for name, path, cpu_limit, memory_limit, enforce_headroom in definitions + ] + report = { + "valid": all(result["valid"] for result in results), + "threshold": args.threshold, + "expected_duration_seconds": args.expected_duration, + "resources": results, + } + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8") + if not report["valid"]: + for result in results: + if not result["valid"]: + print(f"{result['name']}: {result['reason']}") + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmark/verify_server.py b/benchmark/verify_server.py new file mode 100644 index 0000000..9751642 --- /dev/null +++ b/benchmark/verify_server.py @@ -0,0 +1,913 @@ +#!/usr/bin/env python3 +"""Fast correctness gate for the protocol-correct benchmark workload.""" + +from __future__ import annotations + +import argparse +import json +import math +import re +import socket +import urllib.error +import urllib.parse +import urllib.request +from pathlib import Path +from typing import Any + + +PROTOCOL_VERSION = "2024-11-05" +ELIGIBILITY_CONTRACT = "upstream-v2-strict-mcp-v1" +SUPPLEMENTAL_CONTRACT = "adapter-exact-v1" +SUPPORTED_PROTOCOL_VERSIONS = { + "2024-11-05", + "2025-03-26", + "2025-06-18", + "2025-11-25", +} +MCP_RESPONSE_MEDIA_TYPES = {"application/json", "text/event-stream"} +EXPECTED_TOOL_SCHEMAS = { + "search_products": { + "category": {"type": "string"}, + "min_price": {"type": "number"}, + "max_price": {"type": "number"}, + "limit": {"type": "integer"}, + }, + "get_user_cart": {"user_id": {"type": "string"}}, + "checkout": { + "user_id": {"type": "string"}, + "items": { + "type": "array", + "items": { + "type": "object", + "properties": { + "product_id": {"type": "integer"}, + "quantity": {"type": "integer"}, + }, + "required": ["product_id", "quantity"], + }, + }, + }, +} +CANONICAL_TOOL_ARGUMENTS = { + "search_products": { + "category": "Electronics", + "min_price": 50.0, + "max_price": 500.0, + "limit": 10, + }, + "get_user_cart": {"user_id": "user-00001"}, + "checkout": { + "user_id": "user-00001", + "items": [ + {"product_id": 42, "quantity": 2}, + {"product_id": 1337, "quantity": 1}, + ], + }, +} +CANONICAL_POPULAR_IDS = [ + 92857, + 82857, + 72857, + 62857, + 52857, + 42857, + 32857, + 2857, + 22857, + 12857, +] +SCHEMA_ANNOTATION_KEYWORDS = { + "$id", + "$schema", + "$comment", + "title", + "description", + "default", + "examples", + "deprecated", + "readOnly", + "writeOnly", +} +INVALID_JSON_POINTER_ESCAPE = re.compile(r"~(?![01])") +JSON_SCHEMA_2020_12_DIALECTS = { + "https://json-schema.org/draft/2020-12/schema", + "https://json-schema.org/draft/2020-12/schema#", +} + + +def parse_response(body: bytes, media_type: str) -> dict[str, Any] | None: + text = body.decode("utf-8").strip() + if not text: + return None + if media_type == "application/json": + message = json.loads(text) + if not isinstance(message, dict): + raise RuntimeError("a non-batch request returned a batch response") + return message + if media_type != "text/event-stream": + raise RuntimeError(f"unsupported MCP response media type: {media_type!r}") + + messages: list[dict[str, Any]] = [] + for event in text.replace("\r\n", "\n").split("\n\n"): + data = "\n".join( + line[5:].lstrip() + for line in event.splitlines() + if line.startswith("data:") + ) + if not data: + continue + message = json.loads(data) + if not isinstance(message, dict): + raise RuntimeError("an SSE event contained a non-object JSON-RPC message") + messages.append(message) + if len(messages) != 1: + raise RuntimeError( + f"expected exactly one JSON-RPC response, received {len(messages)}" + ) + return messages[0] + + +def response_status(response: Any) -> int: + status = getattr(response, "status", None) + if status is None: + status = response.getcode() + return int(status) + + +def response_media_type(response: Any) -> str: + content_type = response.headers.get("Content-Type", "") + return content_type.split(";", 1)[0].strip().lower() + + +def redis_command(redis_url: str, *arguments: str) -> str | int | None: + """Execute the small RESP subset needed by the workload side-effect gate.""" + parsed = urllib.parse.urlparse(redis_url) + if parsed.scheme != "redis" or not parsed.hostname: + raise RuntimeError(f"unsupported Redis URL: {redis_url!r}") + if parsed.username or parsed.password or parsed.path not in {"", "/", "/0"}: + raise RuntimeError("the benchmark verifier supports only unauthenticated Redis DB 0") + + encoded = [argument.encode("utf-8") for argument in arguments] + request = [f"*{len(encoded)}\r\n".encode("ascii")] + for argument in encoded: + request.extend( + (f"${len(argument)}\r\n".encode("ascii"), argument, b"\r\n") + ) + + with socket.create_connection((parsed.hostname, parsed.port or 6379), timeout=5) as connection: + connection.sendall(b"".join(request)) + with connection.makefile("rb") as response: + prefix = response.read(1) + line = response.readline() + if not prefix or not line.endswith(b"\r\n"): + raise RuntimeError("Redis returned a truncated response") + value = line[:-2] + if prefix == b"+": + return value.decode("utf-8") + if prefix == b"-": + raise RuntimeError(f"Redis command failed: {value.decode('utf-8')}") + if prefix == b":": + return int(value) + if prefix != b"$": + raise RuntimeError(f"unsupported Redis response prefix: {prefix!r}") + length = int(value) + if length == -1: + return None + payload = response.read(length) + terminator = response.read(2) + if len(payload) != length or terminator != b"\r\n": + raise RuntimeError("Redis returned a truncated bulk string") + return payload.decode("utf-8") + + +def redis_workload_state(redis_url: str, user_id: str = "user-00001") -> dict[str, Any]: + rate_value = redis_command(redis_url, "GET", f"bench:ratelimit:{user_id}") + history_length = redis_command(redis_url, "LLEN", f"bench:history:{user_id}") + popularity_score = redis_command( + redis_url, "ZSCORE", "bench:popular", "product:42" + ) + return { + "rate_limit_count": 0 if rate_value is None else int(rate_value), + "history_length": int(history_length), + "product_42_popularity": ( + None if popularity_score is None else float(popularity_score) + ), + } + + +def valid_workload_side_effects( + before: dict[str, Any], after: dict[str, Any] +) -> bool: + before_popularity = before.get("product_42_popularity") + after_popularity = after.get("product_42_popularity") + return ( + isinstance(before.get("rate_limit_count"), int) + and isinstance(after.get("rate_limit_count"), int) + and after["rate_limit_count"] == before["rate_limit_count"] + 1 + and isinstance(before.get("history_length"), int) + and isinstance(after.get("history_length"), int) + and after["history_length"] == before["history_length"] + 1 + and isinstance(before_popularity, (int, float)) + and not isinstance(before_popularity, bool) + and isinstance(after_popularity, (int, float)) + and not isinstance(after_popularity, bool) + and same_number(after_popularity, float(before_popularity) + 1) + ) + + +def exact_keys(value: Any, expected: set[str]) -> bool: + return isinstance(value, dict) and set(value) == expected + + +def same_number(actual: Any, expected: float) -> bool: + return isinstance(actual, (int, float)) and not isinstance(actual, bool) and math.isclose( + float(actual), expected, rel_tol=0.0, abs_tol=1e-9 + ) + + +def expected_search_products() -> list[dict[str, Any]]: + brands = ["Alpha", "Phi", "Cast", "Lambda", "Forge"] + products = [] + for index in range(10): + product_id = 4901 + index * 20 + brand = brands[index % len(brands)] + products.append( + { + "id": product_id, + "sku": f"SKU-{product_id:06d}", + "name": f"{brand} Electronics Item {product_id}", + "price": 50 + index * 0.2, + "rating": 3 if index % 2 == 0 else 1, + "popularity_rank": 0, + } + ) + return products + + +def exact_object(actual: Any, expected: dict[str, Any]) -> bool: + if not exact_keys(actual, set(expected)): + return False + for key, value in expected.items(): + if isinstance(value, (int, float)) and not isinstance(value, bool): + if not same_number(actual[key], float(value)): + return False + elif actual[key] != value: + return False + return True + + +def schema_fragment_matches( + actual: Any, + expected: Any, + keyword: str | None = None, + root_schema: Any | None = None, + instances: list[Any] | None = None, +) -> bool: + if root_schema is None: + root_schema = actual + resolved, actual = resolve_local_schema_reference(actual, root_schema) + if not resolved: + return False + if isinstance(expected, dict): + if not isinstance(actual, dict): + return False + if not schema_extras_are_compatible( + actual, set(expected), root_schema, instances + ): + return False + for name, expected_value in expected.items(): + actual_value = actual.get(name) + if name == "properties" and ( + not isinstance(actual_value, dict) + or set(actual_value) != set(expected_value) + ): + return False + child_instances = instances + if keyword == "properties" and instances is not None: + child_instances = [ + instance[name] + for instance in instances + if isinstance(instance, dict) and name in instance + ] + elif name == "items" and instances is not None: + child_instances = [ + item + for instance in instances + if isinstance(instance, list) + for item in instance + ] + if not schema_fragment_matches( + actual_value, + expected_value, + name, + root_schema, + child_instances, + ): + return False + return True + if isinstance(expected, list): + if not isinstance(actual, list) or len(actual) != len(expected): + return False + if keyword == "required": + return set(actual) == set(expected) + return all( + schema_fragment_matches( + value, + expected[index], + root_schema=root_schema, + ) + for index, value in enumerate(actual) + ) + return actual == expected + + +def schema_extras_are_compatible( + schema: dict[str, Any], + compared_keywords: set[str], + root_schema: Any, + instances: list[Any] | None, +) -> bool: + """Allow only extras that cannot reject any benchmark request.""" + for keyword, value in schema.items(): + if keyword in compared_keywords or keyword in SCHEMA_ANNOTATION_KEYWORDS: + continue + if keyword in {"$defs", "definitions"}: + if not isinstance(value, dict): + return False + continue + if keyword == "format": + dialect = root_schema.get("$schema") if isinstance(root_schema, dict) else None + if ( + not isinstance(value, str) + or (dialect is not None and dialect not in JSON_SCHEMA_2020_12_DIALECTS) + ): + return False + continue + if keyword == "minimum": + if ( + schema.get("type") not in {"integer", "number"} + or isinstance(value, bool) + or not isinstance(value, (int, float)) + or not math.isfinite(float(value)) + or not instances + or any( + isinstance(instance, bool) + or not isinstance(instance, (int, float)) + or not math.isfinite(float(instance)) + or instance < value + for instance in instances + ) + ): + return False + continue + if keyword == "additionalProperties": + if not isinstance(value, bool): + return False + continue + if keyword == "required": + properties = schema.get("properties") + if ( + not isinstance(value, list) + or any(not isinstance(name, str) for name in value) + or len(value) != len(set(value)) + or not isinstance(properties, dict) + or not set(value).issubset(properties) + ): + return False + continue + return False + return True + + +def resolve_local_schema_reference( + value: Any, + root_schema: Any, +) -> tuple[bool, Any]: + """Resolve a chain of local JSON Pointer references without fetching schemas.""" + current = value + seen: set[str] = set() + while isinstance(current, dict) and "$ref" in current: + reference = current.get("$ref") + if ( + not isinstance(reference, str) + or reference in seen + or any( + keyword not in SCHEMA_ANNOTATION_KEYWORDS and keyword != "$ref" + for keyword in current + ) + ): + return False, None + seen.add(reference) + if reference == "#": + current = root_schema + continue + if not reference.startswith("#/"): + return False, None + current = root_schema + for encoded_token in reference[2:].split("/"): + if INVALID_JSON_POINTER_ESCAPE.search(encoded_token): + return False, None + token = encoded_token.replace("~1", "/").replace("~0", "~") + if not isinstance(current, dict) or token not in current: + return False, None + current = current[token] + return True, current + + +def validate_tool_schema(tool: dict[str, Any]) -> bool: + name = tool.get("name") + expected = EXPECTED_TOOL_SCHEMAS.get(name) + arguments = CANONICAL_TOOL_ARGUMENTS.get(name) + schema = tool.get("inputSchema") + properties = schema.get("properties") if isinstance(schema, dict) else None + return ( + expected is not None + and arguments is not None + and isinstance(schema, dict) + and schema.get("type") == "object" + and isinstance(properties, dict) + and set(properties) == set(expected) + and schema_extras_are_compatible( + schema, {"type", "properties"}, schema, [arguments] + ) + and all( + schema_fragment_matches( + properties[property], + property_schema, + root_schema=schema, + instances=[arguments[property]], + ) + for property, property_schema in expected.items() + ) + ) + + +class McpSession: + def __init__(self, url: str, expected_protocol_version: str | None = None) -> None: + self.url = url + self.session_id: str | None = None + self.protocol_version: str | None = None + response = self.post( + { + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": PROTOCOL_VERSION, + "capabilities": {}, + "clientInfo": {"name": "benchmark-verify", "version": "1.0"}, + }, + } + ) + if not response or response.get("jsonrpc") != "2.0" or response.get("id") != 1: + raise RuntimeError(f"initialize returned an invalid JSON-RPC response: {response!r}") + if response.get("error") or not isinstance(response.get("result"), dict): + raise RuntimeError(f"initialize failed: {response!r}") + result = response["result"] + negotiated_version = result.get("protocolVersion") + capabilities = result.get("capabilities") + server_info = result.get("serverInfo") + if negotiated_version not in SUPPORTED_PROTOCOL_VERSIONS: + raise RuntimeError(f"initialize negotiated an unsupported protocol: {response!r}") + if ( + expected_protocol_version is not None + and negotiated_version != expected_protocol_version + ): + raise RuntimeError( + "initialize negotiation changed from " + f"{expected_protocol_version} to {negotiated_version}" + ) + if not isinstance(capabilities, dict) or not isinstance( + capabilities.get("tools"), dict + ): + raise RuntimeError(f"initialize did not advertise tools: {response!r}") + if ( + not isinstance(server_info, dict) + or not isinstance(server_info.get("name"), str) + or not server_info["name"] + or not isinstance(server_info.get("version"), str) + or not server_info["version"] + ): + raise RuntimeError(f"initialize returned invalid serverInfo: {response!r}") + self.protocol_version = negotiated_version + self.post({"jsonrpc": "2.0", "method": "notifications/initialized"}) + + def post(self, payload: dict[str, Any]) -> dict[str, Any] | None: + headers = { + "Content-Type": "application/json", + "Accept": "application/json, text/event-stream", + } + if self.session_id: + headers["Mcp-Session-Id"] = self.session_id + if self.protocol_version: + headers["MCP-Protocol-Version"] = self.protocol_version + request = urllib.request.Request( + self.url, + data=json.dumps(payload).encode("utf-8"), + headers=headers, + method="POST", + ) + with urllib.request.urlopen(request, timeout=30) as response: + response_session_id = response.headers.get("Mcp-Session-Id") + if ( + self.session_id + and response_session_id + and response_session_id != self.session_id + ): + raise RuntimeError("server changed the MCP session identifier") + if response_session_id: + self.session_id = response_session_id + body = response.read() + if "id" not in payload: + if response_status(response) != 202 or body.strip(): + raise RuntimeError( + "notification response must be HTTP 202 with an empty body" + ) + return None + if response_status(response) != 200: + raise RuntimeError( + f"request response must be HTTP 200, got {response_status(response)}" + ) + media_type = response_media_type(response) + if media_type not in MCP_RESPONSE_MEDIA_TYPES: + raise RuntimeError( + f"invalid MCP response Content-Type: {media_type!r}" + ) + message = parse_response(body, media_type) + if ( + not message + or message.get("jsonrpc") != "2.0" + or message.get("id") != payload["id"] + ): + raise RuntimeError(f"invalid JSON-RPC response: {message!r}") + return message + + def close(self) -> None: + if not self.session_id: + return + headers = {"Mcp-Session-Id": self.session_id} + if self.protocol_version: + headers["MCP-Protocol-Version"] = self.protocol_version + request = urllib.request.Request(self.url, headers=headers, method="DELETE") + try: + with urllib.request.urlopen(request, timeout=5): + pass + except urllib.error.HTTPError as error: + if error.code != 405: + raise + + def __enter__(self) -> "McpSession": + return self + + def __exit__(self, *_: object) -> None: + self.close() + + +def call_tool( + url: str, + name: str, + arguments: dict[str, Any], + expected_protocol_version: str, +) -> dict[str, Any]: + with McpSession(url, expected_protocol_version) as session: + response = session.post( + { + "jsonrpc": "2.0", + "id": 2, + "method": "tools/call", + "params": {"name": name, "arguments": arguments}, + } + ) + if not response or response.get("error") or "result" not in response: + raise RuntimeError(f"{name} failed: {response!r}") + result = response["result"] + if not isinstance(result, dict): + raise RuntimeError(f"{name} returned an invalid result: {response!r}") + if result.get("isError") is True: + raise RuntimeError(f"{name} returned isError: {response!r}") + content = result.get("content", []) + if ( + not isinstance(content, list) + or not content + or not isinstance(content[0], dict) + or content[0].get("type") != "text" + or not isinstance(content[0].get("text"), str) + ): + raise RuntimeError(f"{name} returned no text content: {response!r}") + value = json.loads(content[0]["text"]) + if not isinstance(value, dict): + raise RuntimeError(f"{name} text content is not a JSON object: {response!r}") + return value + + +def valid_upstream_search(value: dict[str, Any]) -> bool: + return ( + value.get("total_found") == 2251 + and isinstance(value.get("products"), list) + and len(value["products"]) == 10 + and isinstance(value.get("top10_popular_ids"), list) + and len(value["top10_popular_ids"]) == 10 + ) + + +def valid_upstream_cart(value: dict[str, Any], user_id: str) -> bool: + cart = value.get("cart") + return ( + value.get("user_id") == user_id + and isinstance(cart, dict) + and isinstance(cart.get("items"), list) + and len(cart["items"]) >= 1 + and isinstance(value.get("recent_history"), list) + and len(value["recent_history"]) == 5 + ) + + +def valid_upstream_checkout(value: dict[str, Any], user_id: str) -> bool: + total = value.get("total") + rate_limit_count = value.get("rate_limit_count") + return ( + value.get("user_id") == user_id + and value.get("status") == "confirmed" + and isinstance(total, (int, float)) + and not isinstance(total, bool) + and total > 0 + and value.get("items_count") == 2 + and isinstance(rate_limit_count, (int, float)) + and not isinstance(rate_limit_count, bool) + ) + + +def valid_exact_search( + value: dict[str, Any], expected_server_type: str | None +) -> bool: + expected_products = expected_search_products() + return ( + exact_keys( + value, + {"category", "total_found", "products", "top10_popular_ids", "server_type"}, + ) + and value["category"] == "Electronics" + and value["total_found"] == 2251 + and isinstance(value["products"], list) + and len(value["products"]) == len(expected_products) + and all( + exact_object(product, expected_products[index]) + for index, product in enumerate(value["products"]) + ) + and value["top10_popular_ids"] == CANONICAL_POPULAR_IDS + and ( + expected_server_type is None + or value["server_type"] == expected_server_type + ) + ) + + +def valid_exact_cart( + value: dict[str, Any], expected_server_type: str | None +) -> bool: + expected_items = [{"product_id": 8, "qty": 2}, {"product_id": 14, "qty": 2}] + expected_history = [ + { + "order_id": f"ORD-00001-{entry:02d}", + "product_id": entry * 7 + 1, + "qty": 1 + entry % 3, + "price": round((entry * 13 + 1) / 100, 2), + "ts": 1740000000 + 86400 + entry * 3600, + } + for entry in range(1, 6) + ] + cart = value.get("cart") + return ( + exact_keys(value, {"user_id", "cart", "recent_history", "server_type"}) + and value["user_id"] == "user-00001" + and exact_keys(cart, {"items", "item_count", "estimated_total"}) + and cart["items"] == expected_items + and cart["item_count"] == 2 + and same_number(cart["estimated_total"], 44.04) + and isinstance(value["recent_history"], list) + and len(value["recent_history"]) == len(expected_history) + and all( + exact_object(history, expected_history[index]) + for index, history in enumerate(value["recent_history"]) + ) + and ( + expected_server_type is None + or value["server_type"] == expected_server_type + ) + ) + + +def valid_exact_checkout( + value: dict[str, Any], expected_server_type: str | None +) -> bool: + return ( + exact_keys( + value, + { + "order_id", + "user_id", + "total", + "items_count", + "rate_limit_count", + "status", + "server_type", + }, + ) + and value["order_id"] == "ORD-user00001-2" + and value["user_id"] == "user-00001" + and same_number(value["total"], 24.63) + and value["items_count"] == 2 + and value["rate_limit_count"] == 1 + and value["status"] == "confirmed" + and ( + expected_server_type is None + or value["server_type"] == expected_server_type + ) + ) + + +def verify( + url: str, + expected_protocol_version: str | None = None, + expected_server_type: str | None = None, + redis_url: str | None = None, +) -> dict[str, Any]: + with McpSession(url, expected_protocol_version) as session: + negotiated_protocol_version = session.protocol_version + response = session.post( + {"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}} + ) + if not response or response.get("error") or not isinstance(response.get("result"), dict): + raise RuntimeError(f"tools/list failed: {response!r}") + tools = response["result"].get("tools") + if not isinstance(tools, list) or not all(isinstance(tool, dict) for tool in tools): + raise RuntimeError(f"tools/list returned an invalid tools array: {response!r}") + names = [tool.get("name") for tool in tools] + required = set(EXPECTED_TOOL_SCHEMAS) + if len(tools) != len(required) or set(names) != required: + raise RuntimeError(f"tools/list returned the wrong workload tools: {names!r}") + schema_objects = { + str(tool["name"]): isinstance(tool.get("inputSchema"), dict) for tool in tools + } + if not all(schema_objects.values()): + raise RuntimeError(f"tools/list returned an invalid input schema object: {tools!r}") + exact_schema_checks = { + str(tool["name"]): validate_tool_schema(tool) for tool in tools + } + + assert negotiated_protocol_version is not None + search = call_tool( + url, + "search_products", + {"category": "Electronics", "min_price": 50, "max_price": 500, "limit": 10}, + negotiated_protocol_version, + ) + if not valid_upstream_search(search): + raise RuntimeError(f"invalid search_products result: {search!r}") + + cart = call_tool( + url, + "get_user_cart", + {"user_id": "user-00001"}, + negotiated_protocol_version, + ) + if not valid_upstream_cart(cart, "user-00001"): + raise RuntimeError(f"invalid get_user_cart result: {cart!r}") + + redis_before = redis_workload_state(redis_url) if redis_url else None + checkout = call_tool( + url, + "checkout", + { + "user_id": "user-00001", + "items": [ + {"product_id": 42, "quantity": 2}, + {"product_id": 1337, "quantity": 1}, + ], + }, + negotiated_protocol_version, + ) + if not valid_upstream_checkout(checkout, "user-00001"): + raise RuntimeError(f"invalid checkout result: {checkout!r}") + redis_after = redis_workload_state(redis_url) if redis_url else None + side_effects_valid = ( + redis_before is not None + and redis_after is not None + and valid_workload_side_effects(redis_before, redis_after) + ) + if redis_url and not side_effects_valid: + raise RuntimeError( + "checkout did not produce the required Redis side effects: " + f"before={redis_before!r}, after={redis_after!r}" + ) + + supplemental_checks = { + "exact_tool_schemas": all(exact_schema_checks.values()), + "exact_search_fixture": valid_exact_search(search, expected_server_type), + "exact_cart_fixture": valid_exact_cart(cart, expected_server_type), + "exact_checkout_fixture": valid_exact_checkout( + checkout, expected_server_type + ), + } + return { + "negotiated_protocol_version": negotiated_protocol_version, + "eligibility_contract": ELIGIBILITY_CONTRACT, + "eligibility_valid": True, + "eligibility_checks": { + "exact_tool_names": True, + "input_schemas_are_objects": schema_objects, + "search_products": True, + "get_user_cart": True, + "checkout": True, + "redis_side_effects": side_effects_valid if redis_url else None, + }, + "supplemental_validation": { + "contract": SUPPLEMENTAL_CONTRACT, + "valid": all(supplemental_checks.values()), + "checks": supplemental_checks, + "tool_schema_checks": exact_schema_checks, + "observations": { + "checkout_rate_limit_count": checkout.get("rate_limit_count"), + "redis_before": redis_before, + "redis_after": redis_after, + }, + }, + } + + +def initialize_only(url: str, expected_protocol_version: str | None = None) -> str: + with McpSession(url, expected_protocol_version) as session: + assert session.protocol_version is not None + return session.protocol_version + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("url") + parser.add_argument("--name", default="server") + parser.add_argument("--expected-protocol-version") + parser.add_argument("--expected-server-type") + parser.add_argument( + "--eligibility-contract", + choices=[ELIGIBILITY_CONTRACT], + default=ELIGIBILITY_CONTRACT, + ) + parser.add_argument("--redis-url") + parser.add_argument("--require-supplemental", action="store_true") + parser.add_argument("--output", type=Path) + parser.add_argument("--initialize-only", action="store_true") + args = parser.parse_args() + if args.initialize_only: + negotiated_protocol_version = initialize_only( + args.url, args.expected_protocol_version + ) + result = { + "schema_version": 2, + "server": args.name, + "requested_protocol_version": PROTOCOL_VERSION, + "negotiated_protocol_version": negotiated_protocol_version, + "eligibility_contract": ELIGIBILITY_CONTRACT, + "eligibility_valid": None, + "supplemental_validation": None, + } + else: + result = verify( + args.url, + args.expected_protocol_version, + args.expected_server_type, + args.redis_url, + ) + result.update( + { + "schema_version": 2, + "server": args.name, + "requested_protocol_version": PROTOCOL_VERSION, + "expected_server_type": args.expected_server_type, + } + ) + result["supplemental_validation"]["required"] = ( + args.require_supplemental + ) + if args.output: + args.output.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8") + if ( + not args.initialize_only + and args.require_supplemental + and not result["supplemental_validation"]["valid"] + ): + raise RuntimeError( + f"required {SUPPLEMENTAL_CONTRACT} validation failed: " + f"{result['supplemental_validation']['checks']!r}" + ) + suffix = ( + "initialize lifecycle passed" + if args.initialize_only + else f"{ELIGIBILITY_CONTRACT} passed" + ) + print( + f"{args.name}: protocol {result['negotiated_protocol_version']}; {suffix}" + ) + + +if __name__ == "__main__": + main() diff --git a/conformance/CMakeLists.txt b/conformance/CMakeLists.txt new file mode 100644 index 0000000..6f50650 --- /dev/null +++ b/conformance/CMakeLists.txt @@ -0,0 +1,5 @@ +add_executable(mcp-conformance-everything-server everything_server.cpp) +target_link_libraries(mcp-conformance-everything-server PRIVATE mcp-cpp-sdk) + +add_executable(mcp-conformance-everything-client everything_client.cpp) +target_link_libraries(mcp-conformance-everything-client PRIVATE mcp-cpp-sdk) diff --git a/conformance/README.md b/conformance/README.md new file mode 100644 index 0000000..94cc3aa --- /dev/null +++ b/conformance/README.md @@ -0,0 +1,22 @@ +# Conformance fixtures + +These fixtures integrate the official MCP conformance runner without claiming +an official SDK tier. The runner is locked to version `0.1.16`, the tested +protocol revision is `2025-11-25`, and the full Tier-required server `active` +and client `core` suites run on every conformance CI job. + +```bash +python scripts/build.py --conformance +npm ci --prefix conformance/runner +bash conformance/run.sh build/release build/conformance-results +``` + +`expected-failures.yml` is a regression baseline, not an exclusion list. Every +scenario still runs. A new failure, or an expected failure that starts passing +without its baseline entry being removed, fails the command. Result directories +contain the individual scenario checks, fixture logs, runner logs, source +revision, and an aggregate scenario-level summary. A green job means the +baseline did not drift; it does not mean that an SDK tier has been achieved. + +The baseline must be regenerated only after reviewing the complete runner +output. Updating it to hide an uninvestigated regression is not acceptable. diff --git a/conformance/everything_client.cpp b/conformance/everything_client.cpp new file mode 100644 index 0000000..04d19a7 --- /dev/null +++ b/conformance/everything_client.cpp @@ -0,0 +1,432 @@ +// Scenario-driven client fixture for the official MCP conformance runner. + +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { + +struct EmptyArguments {}; + +void to_json(nlohmann::json& json, const EmptyArguments&) { json = nlohmann::json::object(); } + +void from_json(const nlohmann::json&, EmptyArguments&) {} + +struct AddNumbersArguments { + int a; + int b; +}; + +void to_json(nlohmann::json& json, const AddNumbersArguments& arguments) { + json = {{"a", arguments.a}, {"b", arguments.b}}; +} + +void from_json(const nlohmann::json& json, AddNumbersArguments& arguments) { + json.at("a").get_to(arguments.a); + json.at("b").get_to(arguments.b); +} + +std::string required_environment(std::string_view name) { + const auto* value = std::getenv(std::string(name).c_str()); + if (value == nullptr || *value == '\0') { + throw std::runtime_error(std::string(name) + " is not set"); + } + return value; +} + +std::optional optional_environment(std::string_view name) { + const auto* value = std::getenv(std::string(name).c_str()); + if (value == nullptr || *value == '\0') { + return std::nullopt; + } + return std::string(value); +} + +nlohmann::json apply_schema_defaults(const nlohmann::json& request) { + nlohmann::json content = nlohmann::json::object(); + if (!request.contains("requestedSchema")) { + return content; + } + + const auto& schema = request.at("requestedSchema"); + if (!schema.contains("properties") || !schema.at("properties").is_object()) { + return content; + } + + for (const auto& [name, property] : schema.at("properties").items()) { + if (property.is_object() && property.contains("default")) { + content[name] = property.at("default"); + } + } + return content; +} + +void install_reverse_rpc_handlers(mcp::Client& client) { + client.on_request( + "elicitation/create", [](const nlohmann::json& request) -> mcp::Task { + nlohmann::json result = {{"action", "accept"}, {"content", apply_schema_defaults(request)}}; + co_return result; + }); + + client.on_request("sampling/createMessage", [](const nlohmann::json&) -> mcp::Task { + nlohmann::json result = { + {"role", "assistant"}, + {"content", {{"type", "text"}, {"text", "Conformance sample response"}}}, + {"model", "mcp-cpp-sdk-conformance-model"}, + {"stopReason", "endTurn"}}; + co_return result; + }); +} + +mcp::Task connect_conformance_client(mcp::Client& client) { + mcp::Implementation implementation; + implementation.name = "mcp-cpp-sdk-conformance-client"; + implementation.version = "0.2.0"; + + mcp::ClientCapabilities capabilities; + mcp::ClientCapabilities::ElicitationCapability elicitation; + elicitation.form = nlohmann::json::object(); + capabilities.elicitation = std::move(elicitation); + capabilities.sampling = mcp::ClientCapabilities::SamplingCapability{}; + capabilities.roots = mcp::ClientCapabilities::RootsCapability{false}; + + return client.connect(implementation, capabilities); +} + +mcp::Task run_initialize(mcp::Client& client) { + (void)co_await connect_conformance_client(client); + (void)co_await client.list_tools(); + client.close(); +} + +mcp::Task run_tools_call(mcp::Client& client) { + (void)co_await connect_conformance_client(client); + (void)co_await client.list_tools(); + (void)co_await client.call_tool("add_numbers", AddNumbersArguments{5, 3}); + client.close(); +} + +mcp::Task run_elicitation_defaults(mcp::Client& client) { + (void)co_await connect_conformance_client(client); + (void)co_await client.call_tool("test_client_elicitation_defaults", EmptyArguments{}); + client.close(); +} + +mcp::Task run_sse_retry(mcp::Client& client) { + (void)co_await connect_conformance_client(client); + (void)co_await client.call_tool("test_reconnection", EmptyArguments{}); + client.close(); +} + +mcp::Task run_scope_step_up(mcp::Client& client) { + (void)co_await connect_conformance_client(client); + (void)co_await client.list_tools(); + // The step-up fixture returns its 403 broader-scope challenge only for tools/call, so the + // escalation leg exists only if the client actually calls a tool after the initial grant. + (void)co_await client.call_tool("test-tool", EmptyArguments{}); + client.close(); +} + +mcp::Task run_scenario(mcp::Client& client, std::string_view scenario) { + if (scenario == "initialize") { + return run_initialize(client); + } + if (scenario == "tools_call") { + return run_tools_call(client); + } + if (scenario == "elicitation-sep1034-client-defaults") { + return run_elicitation_defaults(client); + } + if (scenario == "sse-retry") { + return run_sse_retry(client); + } + if (scenario == "auth/scope-step-up") { + // Kept narrow on purpose: other auth fixtures use server builders that never register + // "test-tool", and only scope-retry-limit tolerates a client error. + return run_scope_step_up(client); + } + if (scenario.starts_with("auth/")) { + // Every other authorization scenario drives the same flow: one authenticated MCP request. + // What differs between them is what the fixture's servers advertise, which the SDK reacts + // to on its own. + return run_initialize(client); + } + throw std::runtime_error("unsupported conformance client scenario: " + std::string(scenario)); +} + +// --- Authorization wiring ------------------------------------------------------------------- +// +// This fixture is a host application: the SDK never opens a browser and never binds a listener, so +// carrying the user agent to the authorization endpoint and collecting the redirect is work the +// application does. Here that means one plain HTTP GET whose redirect is deliberately not followed, +// because the redirect target *is* the authorization response. + +/// The client ID metadata document this fixture publishes. It is never fetched during the flow; +/// the URL is the client identifier itself. +constexpr std::string_view g_client_metadata_url = + "https://conformance-test.local/client-metadata.json"; + +/// Redirect URI registered for the flow. The authorization server only echoes it back in a +/// `Location` header, so nothing ever connects to it. +constexpr std::string_view g_redirect_uri = "http://127.0.0.1:8080/callback"; + +struct RedirectLeg { + std::string host; + std::string port; + std::string target; + std::string location; +}; + +/// Split an `http://host[:port]/path` URL into the pieces a request needs. +RedirectLeg split_http_url(const std::string& url) { + constexpr std::string_view scheme = "http://"; + if (!url.starts_with(scheme)) { + throw std::runtime_error("Authorization URL is not plain HTTP: " + url); + } + auto rest = url.substr(scheme.size()); + const auto path_start = rest.find('/'); + + RedirectLeg leg; + leg.target = path_start == std::string::npos ? "/" : rest.substr(path_start); + auto authority = path_start == std::string::npos ? rest : rest.substr(0, path_start); + + const auto colon = authority.find(':'); + if (colon == std::string::npos) { + leg.host = std::move(authority); + leg.port = "80"; + } else { + leg.host = authority.substr(0, colon); + leg.port = authority.substr(colon + 1); + } + return leg; +} + +/// True for an origin served on this machine's loopback interface. +bool is_loopback_origin(const std::string& origin) { + constexpr std::string_view scheme = "http://"; + if (!origin.starts_with(scheme)) { + return false; + } + auto authority = origin.substr(scheme.size()); + if (authority.starts_with("[")) { + const auto closing = authority.find(']'); + if (closing == std::string::npos) { + return false; + } + authority = authority.substr(1, closing - 1); + } else if (const auto colon = authority.find(':'); colon != std::string::npos) { + authority = authority.substr(0, colon); + } + if (authority == "localhost") { + return true; + } + boost::system::error_code error; + const auto address = boost::asio::ip::make_address(authority, error); + return !error && address.is_loopback(); +} + +mcp::Task run_redirect_leg(boost::asio::any_io_executor executor, + std::shared_ptr leg) { + namespace beast = boost::beast; + namespace http = beast::http; + + boost::asio::ip::tcp::resolver resolver(executor); + const auto endpoints = + co_await resolver.async_resolve(leg->host, leg->port, boost::asio::use_awaitable); + + beast::tcp_stream stream(executor); + stream.expires_after(std::chrono::seconds(10)); + co_await stream.async_connect(endpoints, boost::asio::use_awaitable); + + http::request request(http::verb::get, leg->target, 11); + request.set(http::field::host, leg->host); + co_await http::async_write(stream, request, boost::asio::use_awaitable); + + beast::flat_buffer buffer; + http::response response; + co_await http::async_read(stream, buffer, response, boost::asio::use_awaitable); + + boost::system::error_code ignored; + (void)stream.socket().shutdown(boost::asio::ip::tcp::socket::shutdown_both, ignored); + + const auto location = response.find(http::field::location); + if (location == response.end()) { + throw std::runtime_error("Authorization endpoint returned no redirect to the redirect URI"); + } + leg->location = std::string(location->value()); + co_return mcp::auth::parse_authorization_response(leg->location); +} + +mcp::auth::AuthorizationCallback make_authorization_callback(boost::asio::any_io_executor executor) { + return [executor](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + return run_redirect_leg( + executor, std::make_shared(split_http_url(request.authorization_url))); + }; +} + +/// The outbound-request policy every fixture request runs under. +/// +/// The narrow loopback opt-out lives here and only here: the fixture's servers speak plain HTTP on +/// ephemeral loopback ports, which is not production-ready OAuth. The authorization server's origin +/// is only learned at run time from the protected resource's metadata, so the allow list states the +/// rule rather than an enumeration. +mcp::auth::MetadataFetchPolicy make_fetch_policy(const std::string& server_url) { + mcp::auth::MetadataFetchPolicy policy; + policy.allow_plain_http_loopback = true; + policy.allowed_origins.push_back(mcp::auth::metadata_url_origin(server_url)); + policy.origin_allowance = is_loopback_origin; + return policy; +} + +/// The credentials the runner hands the fixture out of band, if it handed it any. +std::optional conformance_context() { + const auto context = optional_environment("MCP_CONFORMANCE_CONTEXT"); + if (!context) { + return std::nullopt; + } + auto parsed = nlohmann::json::parse(*context, nullptr, false); + if (!parsed.is_object()) { + return std::nullopt; + } + return parsed; +} + +/// Resolve the issuer the runner's injected credentials are bound to. +/// +/// The SDK refuses to present a `client_secret` that names no issuer, because on a real deployment +/// the authorization server is named by a document the client did not write. The fixture is the one +/// case where that document is trustworthy: the runner starts both the MCP server and the +/// authorization server for the scenario, and hands over the credentials it minted for that pair. +/// The runner does not put the issuer in `MCP_CONFORMANCE_CONTEXT` (conformance 0.1.16 sends only +/// `client_id` and `client_secret`), so this reads an `issuer` key when a future runner supplies +/// one and otherwise performs the protected-resource discovery step itself, under the same fetch +/// policy the SDK will use, and binds to the authorization server that document names. +/// +/// This is a deliberate fixture-only decision and must not be copied into an application: an +/// application that binds its secret to whatever issuer the protected-resource document names has +/// re-created the very misbinding the SDK refusal exists to prevent. An application knows its own +/// issuer out of band, because that is where it registered. +std::string resolve_injected_issuer(boost::asio::io_context& io_context, const std::string& server_url, + const nlohmann::json& context) { + if (context.contains("issuer") && context.at("issuer").is_string()) { + return context.at("issuer").get(); + } + + auto http_client = std::make_shared(io_context.get_executor()); + http_client->set_metadata_policy(make_fetch_policy(server_url)); + auto discovery = std::make_shared(http_client); + + auto resolved = boost::asio::co_spawn( + io_context, + [discovery, server_url]() -> mcp::Task { + const auto metadata = co_await discovery->discover_protected_resource(server_url); + if (metadata.authorization_servers.empty()) { + throw std::runtime_error( + "Protected resource metadata listed no authorization server to bind the " + "runner-supplied credentials to"); + } + co_return metadata.authorization_servers.front(); + }, + boost::asio::use_future); + + io_context.run(); + io_context.restart(); + return resolved.get(); +} + +mcp::auth::OAuthAuthorizationConfig make_authorization_config(boost::asio::io_context& io_context, + const std::string& server_url) { + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = server_url; + config.redirect_uri = std::string(g_redirect_uri); + config.client_identity.client_metadata_url = std::string(g_client_metadata_url); + config.client_identity.metadata.client_name = "mcp-cpp-sdk-conformance-client"; + config.credential_store = std::make_shared(); + config.policy = make_fetch_policy(server_url); + + // Credentials the runner hands the fixture out of band. When they are present the SDK presents + // them and never registers; when they are absent it chooses between the metadata document and + // registration on what the authorization server advertises. + if (const auto context = conformance_context(); context && context->contains("client_id")) { + mcp::auth::OAuthClientInformation injected; + injected.client_id = context->at("client_id").get(); + if (context->contains("client_secret")) { + injected.client_secret = context->at("client_secret").get(); + // A secret must name the issuer it belongs to or the SDK refuses to present it. + injected.issuer = resolve_injected_issuer(io_context, server_url, *context); + } + injected.source = mcp::auth::ClientIdentitySource::pre_registered; + config.client_identity.pre_registered = std::move(injected); + } + + return config; +} + +} // namespace + +int main(int argc, char** argv) { + if (argc < 2) { + std::cerr << "Usage: mcp-conformance-everything-client \n"; + return EXIT_FAILURE; + } + + try { + auto scenario = required_environment("MCP_CONFORMANCE_SCENARIO"); + std::string server_url = argv[argc - 1]; + + boost::asio::io_context io_context; + auto executor = io_context.get_executor(); + + // Built before the transport exists, because binding runner-supplied credentials to their + // issuer may need a discovery round trip that drives this io_context to completion first. + std::optional authorization_config; + if (std::string_view(scenario).starts_with("auth/")) { + authorization_config = make_authorization_config(io_context, server_url); + } + + std::shared_ptr transport = + std::make_shared(executor, server_url); + + if (authorization_config) { + auto manager = std::make_shared( + executor, std::make_shared(), + std::move(*authorization_config), make_authorization_callback(executor)); + transport = std::make_shared(transport, manager); + } + + mcp::Client client(transport, executor); + install_reverse_rpc_handlers(client); + + auto completion = + boost::asio::co_spawn(io_context, run_scenario(client, scenario), boost::asio::use_future); + io_context.run(); + completion.get(); + return EXIT_SUCCESS; + } catch (const std::exception& error) { + std::cerr << "Conformance client failed: " << error.what() << '\n'; + return EXIT_FAILURE; + } +} diff --git a/conformance/everything_server.cpp b/conformance/everything_server.cpp new file mode 100644 index 0000000..db6f667 --- /dev/null +++ b/conformance/everything_server.cpp @@ -0,0 +1,479 @@ +// Project fixture for the official MCP conformance runner's 2025-11-25 baseline. + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { + +constexpr std::string_view g_test_image_base64 = + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8DwHwAFBQIAX8jx0gAAAABJRU5ErkJggg=="; +constexpr std::string_view g_test_audio_base64 = + "UklGRiYAAABXQVZFZm10IBAAAAABAAEAQB8AAAB9AAACABAAZGF0YQIAAAA="; + +nlohmann::json empty_object_schema() { + return {{"type", "object"}, {"properties", nlohmann::json::object()}}; +} + +mcp::TextContent text_content(std::string text) { + mcp::TextContent content; + content.text = std::move(text); + return content; +} + +mcp::ImageContent image_content() { + mcp::ImageContent content; + content.data = g_test_image_base64; + content.mimeType = "image/png"; + return content; +} + +mcp::AudioContent audio_content() { + mcp::AudioContent content; + content.data = g_test_audio_base64; + content.mimeType = "audio/wav"; + return content; +} + +mcp::EmbeddedResource embedded_text(std::string uri, std::string mime_type, std::string text) { + mcp::TextResourceContents resource; + resource.uri = std::move(uri); + resource.mimeType = std::move(mime_type); + resource.text = std::move(text); + + mcp::EmbeddedResource content; + content.resource = std::move(resource); + return content; +} + +mcp::PromptMessage prompt_text(std::string text) { + mcp::PromptMessage message; + message.role = mcp::Role::eUser; + message.content = text_content(std::move(text)); + return message; +} + +unsigned short parse_port(int argc, char** argv) { + constexpr unsigned short default_port = 3001; + if (argc < 2) { + return default_port; + } + + unsigned int value = 0; + const std::string_view argument(argv[1]); + const auto [end, error] = + std::from_chars(argument.data(), argument.data() + argument.size(), value); + if (error != std::errc{} || end != argument.data() + argument.size() || value == 0 || + value > 65535) { + throw std::invalid_argument("port must be an integer in the range 1..65535"); + } + return static_cast(value); +} + +void register_content_tools(mcp::Server& server) { + const auto empty_schema = empty_object_schema(); + + server.add_tool( + "test_simple_text", "Return text content", empty_schema, [](const nlohmann::json&) { + return mcp::make_tool_text_result("This is a simple text response for testing."); + }); + + server.add_tool( + "test_image_content", "Return image content", empty_schema, [](const nlohmann::json&) { + mcp::CallToolResult result; + result.content.emplace_back(image_content()); + return result; + }); + + server.add_tool( + "test_audio_content", "Return audio content", empty_schema, [](const nlohmann::json&) { + mcp::CallToolResult result; + result.content.emplace_back(audio_content()); + return result; + }); + + server.add_tool( + "test_embedded_resource", "Return an embedded resource", empty_schema, + [](const nlohmann::json&) { + mcp::CallToolResult result; + result.content.emplace_back(embedded_text("test://embedded-resource", "text/plain", + "This is an embedded resource content.")); + return result; + }); + + server.add_tool( + "test_multiple_content_types", "Return mixed content", empty_schema, [](const nlohmann::json&) { + mcp::CallToolResult result; + result.content.emplace_back(text_content("Multiple content types test:")); + result.content.emplace_back(image_content()); + result.content.emplace_back(embedded_text( + "test://mixed-content-resource", "application/json", R"({"test":"data","value":123})")); + return result; + }); + + server.add_tool( + "test_error_handling", "Return a tool-level error", empty_schema, [](const nlohmann::json&) { + return mcp::make_tool_error_result("This tool intentionally returns an error for testing"); + }); + + server.add_tool( + "test_reconnection", "Exercise SSE reconnection", empty_schema, [](const nlohmann::json&) { + return mcp::make_tool_text_result( + "The baseline transport returned the result on the initial stream"); + }); +} + +void register_context_tools(mcp::Server& server) { + const auto empty_schema = empty_object_schema(); + + server.add_tool( + "test_tool_with_logging", "Emit three log messages", empty_schema, + [](mcp::Context& context, nlohmann::json) -> mcp::Task { + co_await context.log_info("Tool execution started"); + co_await context.log_info("Tool processing data"); + co_await context.log_info("Tool execution completed"); + co_return mcp::make_tool_text_result("Tool with logging executed successfully"); + }); + + server.add_tool( + "test_tool_with_progress", "Emit progress notifications", empty_schema, + [](mcp::Context& context, nlohmann::json) -> mcp::Task { + co_await context.report_progress(0, 100, "Completed step 0 of 100"); + co_await context.report_progress(50, 100, "Completed step 50 of 100"); + co_await context.report_progress(100, 100, "Completed step 100 of 100"); + co_return mcp::make_tool_text_result("Progress complete"); + }); + + const nlohmann::json sampling_schema = {{"type", "object"}, + {"properties", {{"prompt", {{"type", "string"}}}}}, + {"required", nlohmann::json::array({"prompt"})}}; + server.add_tool( + "test_sampling", "Request sampling from the client", sampling_schema, + [](mcp::Context& context, nlohmann::json arguments) -> mcp::Task { + mcp::SamplingMessage message; + message.role = mcp::Role::eUser; + message.content = mcp::SamplingMessageContentBlock{ + text_content(arguments.value("prompt", "Test prompt for sampling"))}; + + mcp::CreateMessageRequestParams request; + request.messages.emplace_back(std::move(message)); + request.maxTokens = 100; + (void)co_await context.sample_llm(request); + co_return mcp::make_tool_text_result("LLM response received"); + }); + + const nlohmann::json elicitation_schema = {{"type", "object"}, + {"properties", {{"message", {{"type", "string"}}}}}, + {"required", nlohmann::json::array({"message"})}}; + server.add_tool( + "test_elicitation", "Request form input from the client", elicitation_schema, + [](mcp::Context& context, nlohmann::json arguments) -> mcp::Task { + mcp::ElicitRequestFormParams request; + request.message = arguments.value("message", "Please provide your information"); + request.requestedSchema = { + {"type", "object"}, + {"properties", + {{"username", {{"type", "string"}, {"description", "User's response"}}}, + {"email", {{"type", "string"}, {"description", "User's email address"}}}}}, + {"required", nlohmann::json::array({"username", "email"})}}; + (void)co_await context.elicit(mcp::ElicitRequestParams{std::move(request)}); + co_return mcp::make_tool_text_result("User response received"); + }); +} + +nlohmann::json defaults_schema() { + return {{"type", "object"}, + {"properties", + {{"name", {{"type", "string"}, {"default", "John Doe"}}}, + {"age", {{"type", "integer"}, {"default", 30}}}, + {"score", {{"type", "number"}, {"default", 95.5}}}, + {"status", + {{"type", "string"}, + {"enum", nlohmann::json::array({"active", "inactive", "pending"})}, + {"default", "active"}}}, + {"verified", {{"type", "boolean"}, {"default", true}}}}}, + {"required", nlohmann::json::array()}}; +} + +nlohmann::json enums_schema() { + const auto titled = nlohmann::json::array({{{"const", "value1"}, {"title", "First Option"}}, + {{"const", "value2"}, {"title", "Second Option"}}}); + return { + {"type", "object"}, + {"properties", + {{"untitledSingle", + {{"type", "string"}, {"enum", nlohmann::json::array({"option1", "option2", "option3"})}}}, + {"titledSingle", {{"type", "string"}, {"oneOf", titled}}}, + {"legacyEnum", + {{"type", "string"}, + {"enum", nlohmann::json::array({"opt1", "opt2", "opt3"})}, + {"enumNames", nlohmann::json::array({"Option One", "Option Two", "Option Three"})}}}, + {"untitledMulti", + {{"type", "array"}, + {"items", + {{"type", "string"}, + {"enum", nlohmann::json::array({"option1", "option2", "option3"})}}}}}, + {"titledMulti", {{"type", "array"}, {"items", {{"anyOf", titled}}}}}}}, + {"required", nlohmann::json::array()}}; +} + +void register_elicitation_schema_tool(mcp::Server& server, std::string name, + nlohmann::json requested_schema) { + server.add_tool( + name, "Exercise elicitation schema support", empty_object_schema(), + [schema = std::move(requested_schema)](mcp::Context& context, + nlohmann::json) -> mcp::Task { + mcp::ElicitRequestFormParams request; + request.message = "Conformance elicitation schema test"; + request.requestedSchema = schema; + (void)co_await context.elicit(mcp::ElicitRequestParams{std::move(request)}); + co_return mcp::make_tool_text_result("Elicitation completed"); + }); +} + +void register_schema_tools(mcp::Server& server) { + register_elicitation_schema_tool(server, "test_elicitation_sep1034_defaults", defaults_schema()); + register_elicitation_schema_tool(server, "test_elicitation_sep1330_enums", enums_schema()); + + const nlohmann::json schema_2020_12 = { + {"$schema", "https://json-schema.org/draft/2020-12/schema"}, + {"type", "object"}, + {"$defs", + {{"address", {{"type", "object"}, {"properties", {{"street", {{"type", "string"}}}}}}}}}, + {"properties", {{"address", {{"$ref", "#/$defs/address"}}}}}, + {"additionalProperties", false}}; + server.add_tool( + "json_schema_2020_12_tool", "Tool with JSON Schema 2020-12 features", schema_2020_12, + [](const nlohmann::json&) { return mcp::make_tool_text_result("Schema preserved"); }); +} + +void register_resources(mcp::Server& server) { + mcp::Resource text_resource; + text_resource.uri = "test://static-text"; + text_resource.name = "Static text resource"; + text_resource.description = "Conformance text resource"; + text_resource.mimeType = "text/plain"; + server.add_resource( + text_resource, [](mcp::ReadResourceRequestParams request) { + mcp::ReadResourceResult result; + mcp::TextResourceContents contents; + contents.uri = std::move(request.uri); + contents.mimeType = "text/plain"; + contents.text = "This is the content of the static text resource."; + result.contents.emplace_back(std::move(contents)); + return result; + }); + + mcp::Resource binary_resource; + binary_resource.uri = "test://static-binary"; + binary_resource.name = "Static binary resource"; + binary_resource.description = "Conformance binary resource"; + binary_resource.mimeType = "image/png"; + server.add_resource( + binary_resource, [](mcp::ReadResourceRequestParams request) { + mcp::ReadResourceResult result; + mcp::BlobResourceContents contents; + contents.uri = std::move(request.uri); + contents.mimeType = "image/png"; + contents.blob = g_test_image_base64; + result.contents.emplace_back(std::move(contents)); + return result; + }); + + mcp::Resource watched_resource; + watched_resource.uri = "test://watched-resource"; + watched_resource.name = "Watched resource"; + watched_resource.description = "Resource used by subscription scenarios"; + watched_resource.mimeType = "text/plain"; + server.add_resource( + watched_resource, [](mcp::ReadResourceRequestParams request) { + mcp::ReadResourceResult result; + mcp::TextResourceContents contents; + contents.uri = std::move(request.uri); + contents.mimeType = "text/plain"; + contents.text = "Watched resource content"; + result.contents.emplace_back(std::move(contents)); + return result; + }); + + mcp::ResourceTemplate resource_template; + resource_template.uriTemplate = "test://template/{id}/data"; + resource_template.name = "Parameterized test resource"; + resource_template.description = "Substitutes the id URI segment"; + resource_template.mimeType = "application/json"; + server.add_resource_template( + resource_template, [](mcp::ReadResourceRequestParams request) { + mcp::ReadResourceResult result; + mcp::TextResourceContents contents; + constexpr std::string_view prefix = "test://template/"; + constexpr std::string_view suffix = "/data"; + const std::string uri = std::move(request.uri); + if (!uri.starts_with(prefix) || !uri.ends_with(suffix) || + uri.size() <= prefix.size() + suffix.size()) { + throw std::invalid_argument("unexpected resource-template URI"); + } + const auto id = uri.substr(prefix.size(), uri.size() - prefix.size() - suffix.size()); + contents.uri = uri; + contents.mimeType = "application/json"; + contents.text = + nlohmann::json{{"id", id}, {"templateTest", true}, {"data", "Data for ID: " + id}} + .dump(); + result.contents.emplace_back(std::move(contents)); + return result; + }); +} + +void register_prompts(mcp::Server& server) { + mcp::Prompt simple; + simple.name = "test_simple_prompt"; + simple.description = "A simple conformance prompt"; + server.add_prompt( + simple, [](mcp::GetPromptRequestParams) { + mcp::GetPromptResult result; + result.messages.emplace_back(prompt_text("This is a simple prompt for testing.")); + return result; + }); + + mcp::Prompt with_arguments; + with_arguments.name = "test_prompt_with_arguments"; + with_arguments.description = "A parameterized conformance prompt"; + with_arguments.arguments = + std::vector{{"arg1", "First test argument", true, std::nullopt}, + {"arg2", "Second test argument", true, std::nullopt}}; + server.add_prompt( + with_arguments, [](mcp::GetPromptRequestParams request) { + const auto arguments = request.arguments.value_or(std::map{}); + const auto arg1 = arguments.contains("arg1") ? arguments.at("arg1") : ""; + const auto arg2 = arguments.contains("arg2") ? arguments.at("arg2") : ""; + mcp::GetPromptResult result; + result.messages.emplace_back( + prompt_text("Prompt with arguments: arg1='" + arg1 + "', arg2='" + arg2 + "'")); + return result; + }); + + mcp::Prompt with_resource; + with_resource.name = "test_prompt_with_embedded_resource"; + with_resource.description = "A prompt containing an embedded resource"; + with_resource.arguments = + std::vector{{"resourceUri", "URI to embed", true, std::nullopt}}; + server.add_prompt( + with_resource, [](mcp::GetPromptRequestParams request) { + const auto arguments = request.arguments.value_or(std::map{}); + const auto uri = + arguments.contains("resourceUri") ? arguments.at("resourceUri") : "test://resource"; + mcp::GetPromptResult result; + mcp::PromptMessage resource_message; + resource_message.role = mcp::Role::eUser; + resource_message.content = + embedded_text(uri, "text/plain", "Embedded resource content for testing."); + result.messages.emplace_back(std::move(resource_message)); + result.messages.emplace_back(prompt_text("Please process the embedded resource above.")); + return result; + }); + + mcp::Prompt with_image; + with_image.name = "test_prompt_with_image"; + with_image.description = "A prompt containing image content"; + server.add_prompt( + with_image, [](mcp::GetPromptRequestParams) { + mcp::GetPromptResult result; + mcp::PromptMessage image_message; + image_message.role = mcp::Role::eUser; + image_message.content = image_content(); + result.messages.emplace_back(std::move(image_message)); + result.messages.emplace_back(prompt_text("Please analyze the image above.")); + return result; + }); +} + +void register_completion(mcp::Server& server) { + server.set_completion_provider( + [](const mcp::CompleteParams& params) -> mcp::Task { + mcp::CompleteResult result; + result.completion.values = {params.argument.value + "-completion"}; + result.completion.total = 1; + result.completion.hasMore = false; + co_return result; + }); +} + +void configure_server(mcp::Server& server) { + register_content_tools(server); + register_context_tools(server); + register_schema_tools(server); + register_resources(server); + register_prompts(server); + register_completion(server); +} + +} // namespace + +int main(int argc, char** argv) { + try { + const auto port = parse_port(argc, argv); + + mcp::Implementation implementation; + implementation.name = "mcp-cpp-sdk-conformance-server"; + implementation.version = "0.2.0"; + + mcp::ServerCapabilities capabilities; + capabilities.tools = mcp::ServerCapabilities::ToolsCapability{true}; + capabilities.resources = mcp::ServerCapabilities::ResourcesCapability{true, true}; + capabilities.prompts = mcp::ServerCapabilities::PromptsCapability{true}; + capabilities.logging = nlohmann::json::object(); + capabilities.completions = nlohmann::json::object(); + + auto server_factory = [implementation, capabilities](const boost::asio::any_io_executor&) { + auto server = std::make_unique(implementation, capabilities); + configure_server(*server); + return server; + }; + + boost::asio::io_context io_context; + mcp::StreamableHttpSessionManager manager(io_context.get_executor(), "127.0.0.1", port, + std::move(server_factory)); + manager.set_allowed_origins( + {"http://127.0.0.1:" + std::to_string(port), "http://localhost:" + std::to_string(port)}); + + boost::asio::signal_set signals(io_context, SIGINT, SIGTERM); + signals.async_wait([&manager](const boost::system::error_code& error, int) { + if (!error) { + manager.close(); + } + }); + + std::exception_ptr listen_error; + boost::asio::co_spawn(io_context, manager.listen(), + [&listen_error, &signals](std::exception_ptr error) { + listen_error = std::move(error); + signals.cancel(); + }); + std::cerr << "MCP conformance server listening on http://127.0.0.1:" << port << "/mcp\n"; + io_context.run(); + if (listen_error) { + std::rethrow_exception(listen_error); + } + return EXIT_SUCCESS; + } catch (const std::exception& error) { + std::cerr << "Conformance server failed: " << error.what() << '\n'; + return EXIT_FAILURE; + } +} diff --git a/conformance/expected-failures.yml b/conformance/expected-failures.yml new file mode 100644 index 0000000..fe1c79a --- /dev/null +++ b/conformance/expected-failures.yml @@ -0,0 +1,14 @@ +# Pinned runner: @modelcontextprotocol/conformance 0.1.16 +# Protocol baseline: 2025-11-25 +# +# Every core client scenario is still executed. Unsupported scenarios stay +# visible here so CI catches both new regressions and stale entries after fixes. +client: + - elicitation-sep1034-client-defaults + - sse-retry +server: + - tools-call-sampling + - tools-call-elicitation + - elicitation-sep1034-defaults + - server-sse-multiple-streams + - elicitation-sep1330-enums diff --git a/conformance/run.sh b/conformance/run.sh new file mode 100755 index 0000000..f5c032b --- /dev/null +++ b/conformance/run.sh @@ -0,0 +1,118 @@ +#!/usr/bin/env bash +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +PROJECT_DIR="$(dirname "$SCRIPT_DIR")" +BUILD_DIR="${1:-$PROJECT_DIR/build/release}" +OUTPUT_DIR="${2:-$PROJECT_DIR/build/conformance-results}" +RUNNER="$SCRIPT_DIR/runner/node_modules/.bin/conformance" +SERVER="$BUILD_DIR/conformance/mcp-conformance-everything-server" +CLIENT="$BUILD_DIR/conformance/mcp-conformance-everything-client" +BASELINE="$SCRIPT_DIR/expected-failures.yml" +SPEC_VERSION="2025-11-25" +PORT="${MCP_CONFORMANCE_PORT:-3001}" +SERVER_TIMEOUT="${MCP_CONFORMANCE_SERVER_TIMEOUT:-15m}" +CLIENT_TIMEOUT="${MCP_CONFORMANCE_CLIENT_TIMEOUT:-10m}" + +if [[ ! -x "$RUNNER" ]]; then + echo "Conformance runner is not installed; run: npm ci --prefix conformance/runner" >&2 + exit 2 +fi +runner_version="$("$RUNNER" --version)" +if [[ "$runner_version" != "0.1.16" ]]; then + echo "Unexpected conformance runner version: $runner_version (expected 0.1.16)" >&2 + exit 2 +fi +if [[ ! -x "$SERVER" || ! -x "$CLIENT" ]]; then + echo "Conformance fixtures are missing; run: python scripts/build.py --conformance" >&2 + exit 2 +fi +if [[ -d "$OUTPUT_DIR" ]] && [[ -n "$(find "$OUTPUT_DIR" -mindepth 1 -print -quit)" ]]; then + echo "Conformance output directory must be empty: $OUTPUT_DIR" >&2 + exit 2 +fi + +mkdir -p "$OUTPUT_DIR/server" "$OUTPUT_DIR/client" + +revision="unknown" +if revision_value="$(git -C "$PROJECT_DIR" rev-parse HEAD 2>/dev/null)"; then + revision="$revision_value" +fi +dirty="unknown" +if status_output="$(git -C "$PROJECT_DIR" status --porcelain 2>/dev/null)"; then + if [[ -n "$status_output" ]]; then + dirty="true" + else + dirty="false" + fi +fi +{ + echo "runner=@modelcontextprotocol/conformance@$runner_version" + echo "protocol_version=$SPEC_VERSION" + echo "source_revision=$revision" + echo "source_tree_dirty=$dirty" + echo "started_at=$(date -u +%Y-%m-%dT%H:%M:%SZ)" +} >"$OUTPUT_DIR/metadata.txt" + +server_pid="" +cleanup() { + if [[ -n "$server_pid" ]] && kill -0 "$server_pid" 2>/dev/null; then + kill "$server_pid" 2>/dev/null || true + wait "$server_pid" 2>/dev/null || true + fi +} +trap cleanup EXIT + +"$SERVER" "$PORT" >"$OUTPUT_DIR/server/stdout.txt" 2>"$OUTPUT_DIR/server/stderr.txt" & +server_pid=$! + +ready=0 +for _ in {1..100}; do + if ! kill -0 "$server_pid" 2>/dev/null; then + echo "Conformance server exited before becoming ready" >&2 + exit 1 + fi + if curl --silent --show-error --connect-timeout 1 --max-time 1 --output /dev/null \ + "http://127.0.0.1:$PORT/mcp"; then + ready=1 + break + fi + sleep 0.1 +done +if [[ "$ready" -ne 1 ]]; then + echo "Conformance server did not become ready" >&2 + exit 1 +fi + +set +e +timeout --signal=INT --kill-after=10s "$SERVER_TIMEOUT" "$RUNNER" server \ + --url "http://127.0.0.1:$PORT/mcp" \ + --suite active \ + --spec-version "$SPEC_VERSION" \ + --expected-failures "$BASELINE" \ + --output-dir "$OUTPUT_DIR/server/results" \ + 2>&1 | tee "$OUTPUT_DIR/server/runner.log" +server_status=${PIPESTATUS[0]} + +timeout --signal=INT --kill-after=10s "$CLIENT_TIMEOUT" "$RUNNER" client \ + --command "exec $CLIENT" \ + --suite core \ + --timeout 30000 \ + --spec-version "$SPEC_VERSION" \ + --expected-failures "$BASELINE" \ + --output-dir "$OUTPUT_DIR/client/results" \ + 2>&1 | tee "$OUTPUT_DIR/client/runner.log" +client_status=${PIPESTATUS[0]} +set -e + +python3 "$SCRIPT_DIR/summarize.py" \ + --server-results "$OUTPUT_DIR/server/results" \ + --client-results "$OUTPUT_DIR/client/results" \ + --output-dir "$OUTPUT_DIR" \ + --server-status "$server_status" \ + --client-status "$client_status" + +if [[ "$server_status" -ne 0 || "$client_status" -ne 0 ]]; then + echo "Conformance regression detected (server=$server_status, client=$client_status)" >&2 + exit 1 +fi diff --git a/conformance/run_alpha.sh b/conformance/run_alpha.sh new file mode 100755 index 0000000..12b45ea --- /dev/null +++ b/conformance/run_alpha.sh @@ -0,0 +1,146 @@ +#!/usr/bin/env bash +# INFORMATIONAL-ONLY, NON-DEFAULT, NON-CI-GATING. +# +# Side-by-side alpha conformance harness. Runs the official MCP conformance +# runner's alpha channel (@modelcontextprotocol/conformance@0.2.0-alpha.11, +# which defaults to spec 2026-07-28) against the SAME already-built fixtures +# used by conformance/run.sh, but writes to a DISTINCT results directory and +# is never invoked by CI or any phase-exit script. +# +# conformance/run.sh and its pinned 0.1.16 / spec-2025-11-25 runner are the +# SDK's only tier-evidence instrument and are left completely untouched by +# this script: no shared runner install, no shared baseline file, no shared +# results directory, no shared summarizer (summarize.py's scenario catalog +# and "runner=0.1.16" label are specific to the pinned suite and do not +# apply to the alpha scenario set, which already differs materially). +# +# The alpha runner is fetched on demand via `npx` (network required) pinned +# to an exact version -- nothing is installed into conformance/runner, which +# stays reserved for the pinned 0.1.16 install. +# +# This script's exit code reflects whether it *ran to completion*, not +# whether alpha scenarios passed or failed. Conformance results here are +# never tier evidence. +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +PROJECT_DIR="$(dirname "$SCRIPT_DIR")" +BUILD_DIR="${1:-$PROJECT_DIR/build/release}" +OUTPUT_DIR="${2:-$PROJECT_DIR/build/conformance-results-alpha}" +RUNNER_PACKAGE="@modelcontextprotocol/conformance@0.2.0-alpha.11" +SERVER="$BUILD_DIR/conformance/mcp-conformance-everything-server" +CLIENT="$BUILD_DIR/conformance/mcp-conformance-everything-client" +SPEC_VERSION="2026-07-28" +PORT="${MCP_CONFORMANCE_ALPHA_PORT:-3002}" +SERVER_TIMEOUT="${MCP_CONFORMANCE_SERVER_TIMEOUT:-15m}" +CLIENT_TIMEOUT="${MCP_CONFORMANCE_CLIENT_TIMEOUT:-10m}" + +if ! command -v npx >/dev/null 2>&1; then + echo "npx is not available; the alpha harness fetches the runner on demand and cannot run without it" >&2 + exit 2 +fi + +echo "Resolving pinned alpha runner: $RUNNER_PACKAGE (network access required)..." >&2 +runner_version="$(npx --yes "$RUNNER_PACKAGE" --version)" +if [[ "$runner_version" != "0.2.0-alpha.11" ]]; then + echo "Unexpected alpha conformance runner version: $runner_version (expected 0.2.0-alpha.11)" >&2 + exit 2 +fi +if [[ ! -x "$SERVER" || ! -x "$CLIENT" ]]; then + echo "Conformance fixtures are missing at $BUILD_DIR/conformance;" \ + "this harness runs against the already-built fixtures and does NOT rebuild them." >&2 + echo "Build them first with: python3 scripts/build.py --conformance" >&2 + exit 2 +fi +if [[ -d "$OUTPUT_DIR" ]] && [[ -n "$(find "$OUTPUT_DIR" -mindepth 1 -print -quit)" ]]; then + echo "Alpha conformance output directory must be empty: $OUTPUT_DIR" >&2 + exit 2 +fi + +mkdir -p "$OUTPUT_DIR/server" "$OUTPUT_DIR/client" + +revision="unknown" +if revision_value="$(git -C "$PROJECT_DIR" rev-parse HEAD 2>/dev/null)"; then + revision="$revision_value" +fi +dirty="unknown" +if status_output="$(git -C "$PROJECT_DIR" status --porcelain 2>/dev/null)"; then + if [[ -n "$status_output" ]]; then + dirty="true" + else + dirty="false" + fi +fi +{ + echo "# INFORMATIONAL ONLY -- NEVER TIER EVIDENCE -- NEVER CI-GATING" + echo "runner=@modelcontextprotocol/conformance@$runner_version" + echo "protocol_version=$SPEC_VERSION" + echo "source_revision=$revision" + echo "source_tree_dirty=$dirty" + echo "started_at=$(date -u +%Y-%m-%dT%H:%M:%SZ)" +} >"$OUTPUT_DIR/metadata.txt" + +server_pid="" +cleanup() { + if [[ -n "$server_pid" ]] && kill -0 "$server_pid" 2>/dev/null; then + kill "$server_pid" 2>/dev/null || true + wait "$server_pid" 2>/dev/null || true + fi +} +trap cleanup EXIT + +"$SERVER" "$PORT" >"$OUTPUT_DIR/server/stdout.txt" 2>"$OUTPUT_DIR/server/stderr.txt" & +server_pid=$! + +ready=0 +for _ in {1..100}; do + if ! kill -0 "$server_pid" 2>/dev/null; then + echo "Conformance server exited before becoming ready" >&2 + exit 1 + fi + if curl --silent --show-error --connect-timeout 1 --max-time 1 --output /dev/null \ + "http://127.0.0.1:$PORT/mcp"; then + ready=1 + break + fi + sleep 0.1 +done +if [[ "$ready" -ne 1 ]]; then + echo "Conformance server did not become ready" >&2 + exit 1 +fi + +set +e +timeout --signal=INT --kill-after=10s "$SERVER_TIMEOUT" npx --yes "$RUNNER_PACKAGE" server \ + --url "http://127.0.0.1:$PORT/mcp" \ + --suite active \ + --spec-version "$SPEC_VERSION" \ + --output-dir "$OUTPUT_DIR/server/results" \ + 2>&1 | tee "$OUTPUT_DIR/server/runner.log" +server_status=${PIPESTATUS[0]} + +timeout --signal=INT --kill-after=10s "$CLIENT_TIMEOUT" npx --yes "$RUNNER_PACKAGE" client \ + --command "exec $CLIENT" \ + --suite core \ + --timeout 30000 \ + --spec-version "$SPEC_VERSION" \ + --output-dir "$OUTPUT_DIR/client/results" \ + 2>&1 | tee "$OUTPUT_DIR/client/runner.log" +client_status=${PIPESTATUS[0]} +set -e + +{ + echo "server_runner_exit_status=$server_status" + echo "client_runner_exit_status=$client_status" + echo "finished_at=$(date -u +%Y-%m-%dT%H:%M:%SZ)" +} >>"$OUTPUT_DIR/metadata.txt" + +echo +echo "Alpha conformance run complete (informational only; not tier evidence)." >&2 +echo "server-side runner exit: $server_status, client-side runner exit: $client_status" >&2 +echo "Results: $OUTPUT_DIR" >&2 +echo "Per-scenario pass/fail totals are in $OUTPUT_DIR/server/runner.log and $OUTPUT_DIR/client/runner.log" >&2 + +# Deliberately always exit 0 here: this harness never gates anything. +# A non-zero exit above (missing runner/fixtures/server-not-ready) already returned early. +exit 0 diff --git a/conformance/runner/package-lock.json b/conformance/runner/package-lock.json new file mode 100644 index 0000000..5c46681 --- /dev/null +++ b/conformance/runner/package-lock.json @@ -0,0 +1,1539 @@ +{ + "name": "mcp-cpp-sdk-conformance-runner", + "version": "0.0.0", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "mcp-cpp-sdk-conformance-runner", + "version": "0.0.0", + "dependencies": { + "@modelcontextprotocol/conformance": "0.1.16" + }, + "engines": { + "node": ">=20" + } + }, + "node_modules/@hono/node-server": { + "version": "1.19.14", + "resolved": "https://registry.npmjs.org/@hono/node-server/-/node-server-1.19.14.tgz", + "integrity": + "sha512-GwtvgtXxnWsucXvbQXkRgqksiH2Qed37H9xHZocE5sA3N8O8O8/8FA3uclQXxXVzc9XBZuEOMK7+r02FmSpHtw==", + "license": "MIT", + "engines": { + "node": ">=18.14.1" + }, + "peerDependencies": { + "hono": "^4" + } + }, + "node_modules/@modelcontextprotocol/conformance": { + "version": "0.1.16", + "resolved": + "https://registry.npmjs.org/@modelcontextprotocol/conformance/-/conformance-0.1.16.tgz", + "integrity": + "sha512-GI7qiN0r39/MH2srVUR3AXaEN0YLCro20lIBbnvc1frBhszenxvUifBuTzxeVQVagILfBzCIcnungUOma8OrgA==", + "license": "MIT", + "dependencies": { + "@modelcontextprotocol/sdk": "^1.27.1", + "@octokit/rest": "^22.0.0", + "commander": "^14.0.2", + "eventsource-parser": "^3.0.6", + "express": "^5.1.0", + "jose": "^6.1.2", + "undici": "^7.19.0", + "yaml": "^2.8.2", + "zod": "^4.3.6" + }, + "bin": { + "conformance": "dist/index.js" + } + }, + "node_modules/@modelcontextprotocol/sdk": { + "version": "1.29.0", + "resolved": "https://registry.npmjs.org/@modelcontextprotocol/sdk/-/sdk-1.29.0.tgz", + "integrity": + "sha512-zo37mZA9hJWpULgkRpowewez1y6ML5GsXJPY8FI0tBBCd77HEvza4jDqRKOXgHNn867PVGCyTdzqpz0izu5ZjQ==", + "license": "MIT", + "dependencies": { + "@hono/node-server": "^1.19.9", + "ajv": "^8.17.1", + "ajv-formats": "^3.0.1", + "content-type": "^1.0.5", + "cors": "^2.8.5", + "cross-spawn": "^7.0.5", + "eventsource": "^3.0.2", + "eventsource-parser": "^3.0.0", + "express": "^5.2.1", + "express-rate-limit": "^8.2.1", + "hono": "^4.11.4", + "jose": "^6.1.3", + "json-schema-typed": "^8.0.2", + "pkce-challenge": "^5.0.0", + "raw-body": "^3.0.0", + "zod": "^3.25 || ^4.0", + "zod-to-json-schema": "^3.25.1" + }, + "engines": { + "node": ">=18" + }, + "peerDependencies": { + "@cfworker/json-schema": "^4.1.1", + "zod": "^3.25 || ^4.0" + }, + "peerDependenciesMeta": { + "@cfworker/json-schema": { + "optional": true + }, + "zod": { + "optional": false + } + } + }, + "node_modules/@octokit/auth-token": { + "version": "6.0.0", + "resolved": "https://registry.npmjs.org/@octokit/auth-token/-/auth-token-6.0.0.tgz", + "integrity": + "sha512-P4YJBPdPSpWTQ1NU4XYdvHvXJJDxM6YwpS0FZHRgP7YFkdVxsWcpWGy/NVqlAA7PcPCnMacXlRm1y2PFZRWL/w==", + "license": "MIT", + "engines": { + "node": ">= 20" + } + }, + "node_modules/@octokit/core": { + "version": "7.0.6", + "resolved": "https://registry.npmjs.org/@octokit/core/-/core-7.0.6.tgz", + "integrity": + "sha512-DhGl4xMVFGVIyMwswXeyzdL4uXD5OGILGX5N8Y+f6W7LhC1Ze2poSNrkF/fedpVDHEEZ+PHFW0vL14I+mm8K3Q==", + "license": "MIT", + "dependencies": { + "@octokit/auth-token": "^6.0.0", + "@octokit/graphql": "^9.0.3", + "@octokit/request": "^10.0.6", + "@octokit/request-error": "^7.0.2", + "@octokit/types": "^16.0.0", + "before-after-hook": "^4.0.0", + "universal-user-agent": "^7.0.0" + }, + "engines": { + "node": ">= 20" + } + }, + "node_modules/@octokit/endpoint": { + "version": "11.0.3", + "resolved": "https://registry.npmjs.org/@octokit/endpoint/-/endpoint-11.0.3.tgz", + "integrity": + "sha512-FWFlNxghg4HrXkD3ifYbS/IdL/mDHjh9QcsNyhQjN8dplUoZbejsdpmuqdA76nxj2xoWPs7p8uX2SNr9rYu0Ag==", + "license": "MIT", + "dependencies": { + "@octokit/types": "^16.0.0", + "universal-user-agent": "^7.0.2" + }, + "engines": { + "node": ">= 20" + } + }, + "node_modules/@octokit/graphql": { + "version": "9.0.3", + "resolved": "https://registry.npmjs.org/@octokit/graphql/-/graphql-9.0.3.tgz", + "integrity": + "sha512-grAEuupr/C1rALFnXTv6ZQhFuL1D8G5y8CN04RgrO4FIPMrtm+mcZzFG7dcBm+nq+1ppNixu+Jd78aeJOYxlGA==", + "license": "MIT", + "dependencies": { + "@octokit/request": "^10.0.6", + "@octokit/types": "^16.0.0", + "universal-user-agent": "^7.0.0" + }, + "engines": { + "node": ">= 20" + } + }, + "node_modules/@octokit/openapi-types": { + "version": "27.0.0", + "resolved": "https://registry.npmjs.org/@octokit/openapi-types/-/openapi-types-27.0.0.tgz", + "integrity": + "sha512-whrdktVs1h6gtR+09+QsNk2+FO+49j6ga1c55YZudfEG+oKJVvJLQi3zkOm5JjiUXAagWK2tI2kTGKJ2Ys7MGA==", + "license": "MIT" + }, + "node_modules/@octokit/plugin-paginate-rest": { + "version": "14.0.0", + "resolved": + "https://registry.npmjs.org/@octokit/plugin-paginate-rest/-/plugin-paginate-rest-14.0.0.tgz", + "integrity": + "sha512-fNVRE7ufJiAA3XUrha2omTA39M6IXIc6GIZLvlbsm8QOQCYvpq/LkMNGyFlB1d8hTDzsAXa3OKtybdMAYsV/fw==", + "license": "MIT", + "dependencies": { + "@octokit/types": "^16.0.0" + }, + "engines": { + "node": ">= 20" + }, + "peerDependencies": { + "@octokit/core": ">=6" + } + }, + "node_modules/@octokit/plugin-request-log": { + "version": "6.0.0", + "resolved": + "https://registry.npmjs.org/@octokit/plugin-request-log/-/plugin-request-log-6.0.0.tgz", + "integrity": + "sha512-UkOzeEN3W91/eBq9sPZNQ7sUBvYCqYbrrD8gTbBuGtHEuycE4/awMXcYvx6sVYo7LypPhmQwwpUe4Yyu4QZN5Q==", + "license": "MIT", + "engines": { + "node": ">= 20" + }, + "peerDependencies": { + "@octokit/core": ">=6" + } + }, + "node_modules/@octokit/plugin-rest-endpoint-methods": { + "version": "17.0.0", + "resolved": + "https://registry.npmjs.org/@octokit/plugin-rest-endpoint-methods/-/plugin-rest-endpoint-methods-17.0.0.tgz", + "integrity": + "sha512-B5yCyIlOJFPqUUeiD0cnBJwWJO8lkJs5d8+ze9QDP6SvfiXSz1BF+91+0MeI1d2yxgOhU/O+CvtiZ9jSkHhFAw==", + "license": "MIT", + "dependencies": { + "@octokit/types": "^16.0.0" + }, + "engines": { + "node": ">= 20" + }, + "peerDependencies": { + "@octokit/core": ">=6" + } + }, + "node_modules/@octokit/request": { + "version": "10.0.11", + "resolved": "https://registry.npmjs.org/@octokit/request/-/request-10.0.11.tgz", + "integrity": + "sha512-+s7HUxjfFqOMS9VlIwDffq0MikjSAK0gSpG73W+meAvVAvX4MBrHYTK5Bj3Uot55qFT4gzUtfzE4mGWY4Br8/Q==", + "license": "MIT", + "dependencies": { + "@octokit/endpoint": "^11.0.3", + "@octokit/request-error": "^7.0.2", + "@octokit/types": "^16.0.0", + "content-type": "^2.0.0", + "json-with-bigint": "^3.5.3", + "universal-user-agent": "^7.0.2" + }, + "engines": { + "node": ">= 20" + } + }, + "node_modules/@octokit/request-error": { + "version": "7.1.0", + "resolved": "https://registry.npmjs.org/@octokit/request-error/-/request-error-7.1.0.tgz", + "integrity": + "sha512-KMQIfq5sOPpkQYajXHwnhjCC0slzCNScLHs9JafXc4RAJI+9f+jNDlBNaIMTvazOPLgb4BnlhGJOTbnN0wIjPw==", + "license": "MIT", + "dependencies": { + "@octokit/types": "^16.0.0" + }, + "engines": { + "node": ">= 20" + } + }, + "node_modules/@octokit/request/node_modules/content-type": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/content-type/-/content-type-2.0.0.tgz", + "integrity": + "sha512-j/O/d7GcZCyNl7/hwZAb606rzqkyvaDctLmckbxLzHvFBzTJHuGEdodATcP3yIRoDrLHkIATJuvzbFlp/ki2cQ==", + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/@octokit/rest": { + "version": "22.0.1", + "resolved": "https://registry.npmjs.org/@octokit/rest/-/rest-22.0.1.tgz", + "integrity": + "sha512-Jzbhzl3CEexhnivb1iQ0KJ7s5vvjMWcmRtq5aUsKmKDrRW6z3r84ngmiFKFvpZjpiU/9/S6ITPFRpn5s/3uQJw==", + "license": "MIT", + "dependencies": { + "@octokit/core": "^7.0.6", + "@octokit/plugin-paginate-rest": "^14.0.0", + "@octokit/plugin-request-log": "^6.0.0", + "@octokit/plugin-rest-endpoint-methods": "^17.0.0" + }, + "engines": { + "node": ">= 20" + } + }, + "node_modules/@octokit/types": { + "version": "16.0.0", + "resolved": "https://registry.npmjs.org/@octokit/types/-/types-16.0.0.tgz", + "integrity": + "sha512-sKq+9r1Mm4efXW1FCk7hFSeJo4QKreL/tTbR0rz/qx/r1Oa2VV83LTA/H/MuCOX7uCIJmQVRKBcbmWoySjAnSg==", + "license": "MIT", + "dependencies": { + "@octokit/openapi-types": "^27.0.0" + } + }, + "node_modules/accepts": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/accepts/-/accepts-2.0.0.tgz", + "integrity": + "sha512-5cvg6CtKwfgdmVqY1WIiXKc3Q1bkRqGLi+2W/6ao+6Y7gu/RCwRuAhGEzh5B4KlszSuTLgZYuqFqo5bImjNKng==", + "license": "MIT", + "dependencies": { + "mime-types": "^3.0.0", + "negotiator": "^1.0.0" + }, + "engines": { + "node": ">= 0.6" + } + }, + "node_modules/ajv": { + "version": "8.20.0", + "resolved": "https://registry.npmjs.org/ajv/-/ajv-8.20.0.tgz", + "integrity": + "sha512-Thbli+OlOj+iMPYFBVBfJ3OmCAnaSyNn4M1vz9T6Gka5Jt9ba/HIR56joy65tY6kx/FCF5VXNB819Y7/GUrBGA==", + "license": "MIT", + "dependencies": { + "fast-deep-equal": "^3.1.3", + "fast-uri": "^3.0.1", + "json-schema-traverse": "^1.0.0", + "require-from-string": "^2.0.2" + }, + "funding": { + "type": "github", + "url": "https://github.com/sponsors/epoberezkin" + } + }, + "node_modules/ajv-formats": { + "version": "3.0.1", + "resolved": "https://registry.npmjs.org/ajv-formats/-/ajv-formats-3.0.1.tgz", + "integrity": + "sha512-8iUql50EUR+uUcdRQ3HDqa6EVyo3docL8g5WJ3FNcWmu62IbkGUue/pEyLBW8VGKKucTPgqeks4fIU1DA4yowQ==", + "license": "MIT", + "dependencies": { + "ajv": "^8.0.0" + }, + "peerDependencies": { + "ajv": "^8.0.0" + }, + "peerDependenciesMeta": { + "ajv": { + "optional": true + } + } + }, + "node_modules/before-after-hook": { + "version": "4.0.0", + "resolved": "https://registry.npmjs.org/before-after-hook/-/before-after-hook-4.0.0.tgz", + "integrity": + "sha512-q6tR3RPqIB1pMiTRMFcZwuG5T8vwp+vUvEG0vuI6B+Rikh5BfPp2fQ82c925FOs+b0lcFQ8CFrL+KbilfZFhOQ==", + "license": "Apache-2.0" + }, + "node_modules/body-parser": { + "version": "2.3.0", + "resolved": "https://registry.npmjs.org/body-parser/-/body-parser-2.3.0.tgz", + "integrity": + "sha512-2cGmJupaNgg+QUwVLAucDuWuoMZ6EX9iHDRswZ5lsNYEmwPaRknMPCLZz07yTzVq/83p4o/wzbDZbBrTvGGTIw==", + "license": "MIT", + "dependencies": { + "bytes": "^3.1.2", + "content-type": "^2.0.0", + "debug": "^4.4.3", + "http-errors": "^2.0.1", + "iconv-lite": "^0.7.2", + "on-finished": "^2.4.1", + "qs": "^6.15.2", + "raw-body": "^3.0.2", + "type-is": "^2.1.0" + }, + "engines": { + "node": ">=18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/body-parser/node_modules/content-type": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/content-type/-/content-type-2.0.0.tgz", + "integrity": + "sha512-j/O/d7GcZCyNl7/hwZAb606rzqkyvaDctLmckbxLzHvFBzTJHuGEdodATcP3yIRoDrLHkIATJuvzbFlp/ki2cQ==", + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/bytes": { + "version": "3.1.2", + "resolved": "https://registry.npmjs.org/bytes/-/bytes-3.1.2.tgz", + "integrity": + "sha512-/Nf7TyzTx6S3yRJObOAV7956r8cr2+Oj8AC5dt8wSP3BQAoeX58NoHyCU8P8zGkNXStjTSi6fzO6F0pBdcYbEg==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/call-bind-apply-helpers": { + "version": "1.0.2", + "resolved": + "https://registry.npmjs.org/call-bind-apply-helpers/-/call-bind-apply-helpers-1.0.2.tgz", + "integrity": + "sha512-Sp1ablJ0ivDkSzjcaJdxEunN5/XvksFJ2sMBFfq6x0ryhQV/2b/KwFe21cMpmHtPOSij8K99/wSfoEuTObmuMQ==", + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0", + "function-bind": "^1.1.2" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/call-bound": { + "version": "1.0.4", + "resolved": "https://registry.npmjs.org/call-bound/-/call-bound-1.0.4.tgz", + "integrity": + "sha512-+ys997U96po4Kx/ABpBCqhA9EuxJaQWDQg7295H4hBphv3IZg0boBKuwYpt4YXp6MZ5AmZQnU/tyMTlRpaSejg==", + "license": "MIT", + "dependencies": { + "call-bind-apply-helpers": "^1.0.2", + "get-intrinsic": "^1.3.0" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/commander": { + "version": "14.0.3", + "resolved": "https://registry.npmjs.org/commander/-/commander-14.0.3.tgz", + "integrity": + "sha512-H+y0Jo/T1RZ9qPP4Eh1pkcQcLRglraJaSLoyOtHxu6AapkjWVCy2Sit1QQ4x3Dng8qDlSsZEet7g5Pq06MvTgw==", + "license": "MIT", + "engines": { + "node": ">=20" + } + }, + "node_modules/content-disposition": { + "version": "1.1.0", + "resolved": + "https://registry.npmjs.org/content-disposition/-/content-disposition-1.1.0.tgz", + "integrity": + "sha512-5jRCH9Z/+DRP7rkvY83B+yGIGX96OYdJmzngqnw2SBSxqCFPd0w2km3s5iawpGX8krnwSGmF0FW5Nhr0Hfai3g==", + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/content-type": { + "version": "1.0.5", + "resolved": "https://registry.npmjs.org/content-type/-/content-type-1.0.5.tgz", + "integrity": + "sha512-nTjqfcBFEipKdXCv4YDQWCfmcLZKm81ldF0pAopTvyrFGVbcR6P/VAAd5G7N+0tTr8QqiU0tFadD6FK4NtJwOA==", + "license": "MIT", + "engines": { + "node": ">= 0.6" + } + }, + "node_modules/cookie": { + "version": "0.7.2", + "resolved": "https://registry.npmjs.org/cookie/-/cookie-0.7.2.tgz", + "integrity": + "sha512-yki5XnKuf750l50uGTllt6kKILY4nQ1eNIQatoXEByZ5dWgnKqbnqmTrBE5B4N7lrMJKQ2ytWMiTO2o0v6Ew/w==", + "license": "MIT", + "engines": { + "node": ">= 0.6" + } + }, + "node_modules/cookie-signature": { + "version": "1.2.2", + "resolved": "https://registry.npmjs.org/cookie-signature/-/cookie-signature-1.2.2.tgz", + "integrity": + "sha512-D76uU73ulSXrD1UXF4KE2TMxVVwhsnCgfAyTg9k8P6KGZjlXKrOLe4dJQKI3Bxi5wjesZoFXJWElNWBjPZMbhg==", + "license": "MIT", + "engines": { + "node": ">=6.6.0" + } + }, + "node_modules/cors": { + "version": "2.8.6", + "resolved": "https://registry.npmjs.org/cors/-/cors-2.8.6.tgz", + "integrity": + "sha512-tJtZBBHA6vjIAaF6EnIaq6laBBP9aq/Y3ouVJjEfoHbRBcHBAHYcMh/w8LDrk2PvIMMq8gmopa5D4V8RmbrxGw==", + "license": "MIT", + "dependencies": { + "object-assign": "^4", + "vary": "^1" + }, + "engines": { + "node": ">= 0.10" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/cross-spawn": { + "version": "7.0.6", + "resolved": "https://registry.npmjs.org/cross-spawn/-/cross-spawn-7.0.6.tgz", + "integrity": + "sha512-uV2QOWP2nWzsy2aMp8aRibhi9dlzF5Hgh5SHaB9OiTGEyDTiJJyx0uy51QXdyWbtAHNua4XJzUKca3OzKUd3vA==", + "license": "MIT", + "dependencies": { + "path-key": "^3.1.0", + "shebang-command": "^2.0.0", + "which": "^2.0.1" + }, + "engines": { + "node": ">= 8" + } + }, + "node_modules/debug": { + "version": "4.4.3", + "resolved": "https://registry.npmjs.org/debug/-/debug-4.4.3.tgz", + "integrity": + "sha512-RGwwWnwQvkVfavKVt22FGLw+xYSdzARwm0ru6DhTVA3umU5hZc28V3kO4stgYryrTlLpuvgI9GiijltAjNbcqA==", + "license": "MIT", + "dependencies": { + "ms": "^2.1.3" + }, + "engines": { + "node": ">=6.0" + }, + "peerDependenciesMeta": { + "supports-color": { + "optional": true + } + } + }, + "node_modules/depd": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/depd/-/depd-2.0.0.tgz", + "integrity": + "sha512-g7nH6P6dyDioJogAAGprGpCtVImJhpPk/roCzdb3fIh61/s/nPsfR6onyMwkCAR/OlC3yBC0lESvUoQEAssIrw==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/dunder-proto": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/dunder-proto/-/dunder-proto-1.0.1.tgz", + "integrity": + "sha512-KIN/nDJBQRcXw0MLVhZE9iQHmG68qAVIBg9CqmUYjmQIhgij9U5MFvrqkUL5FbtyyzZuOeOt0zdeRe4UY7ct+A==", + "license": "MIT", + "dependencies": { + "call-bind-apply-helpers": "^1.0.1", + "es-errors": "^1.3.0", + "gopd": "^1.2.0" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/ee-first": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/ee-first/-/ee-first-1.1.1.tgz", + "integrity": + "sha512-WMwm9LhRUo+WUaRN+vRuETqG89IgZphVSNkdFgeb6sS/E4OrDIN7t48CAewSHXc6C8lefD8KKfr5vY61brQlow==", + "license": "MIT" + }, + "node_modules/encodeurl": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/encodeurl/-/encodeurl-2.0.0.tgz", + "integrity": + "sha512-Q0n9HRi4m6JuGIV1eFlmvJB7ZEVxu93IrMyiMsGC0lrMJMWzRgx6WGquyfQgZVb31vhGgXnfmPNNXmxnOkRBrg==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/es-define-property": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/es-define-property/-/es-define-property-1.0.1.tgz", + "integrity": + "sha512-e3nRfgfUZ4rNGL232gUgX06QNyyez04KdjFrF+LTRoOXmrOgFKDg4BCdsjW8EnT69eqdYGmRpJwiPVYNrCaW3g==", + "license": "MIT", + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/es-errors": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/es-errors/-/es-errors-1.3.0.tgz", + "integrity": + "sha512-Zf5H2Kxt2xjTvbJvP2ZWLEICxA6j+hAmMzIlypy4xcBg1vKVnx89Wy0GbS+kf5cwCVFFzdCFh2XSCFNULS6csw==", + "license": "MIT", + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/es-object-atoms": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/es-object-atoms/-/es-object-atoms-1.1.2.tgz", + "integrity": + "sha512-HWcBoN6NileqtSydK2FqHbS/LoDd2pqrnQHLyJzBj4kOp/ky2MWMN694xOfkK8/SnUsW2DH7EfyVlydKCsm1Zw==", + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/escape-html": { + "version": "1.0.3", + "resolved": "https://registry.npmjs.org/escape-html/-/escape-html-1.0.3.tgz", + "integrity": + "sha512-NiSupZ4OeuGwr68lGIeym/ksIZMJodUGOSCZ/FSnTxcrekbvqrgdUxlJOMpijaKZVjAJrWrGs/6Jy8OMuyj9ow==", + "license": "MIT" + }, + "node_modules/etag": { + "version": "1.8.1", + "resolved": "https://registry.npmjs.org/etag/-/etag-1.8.1.tgz", + "integrity": + "sha512-aIL5Fx7mawVa300al2BnEE4iNvo1qETxLrPI/o05L7z6go7fCw1J6EQmbK4FmJ2AS7kgVF/KEZWufBfdClMcPg==", + "license": "MIT", + "engines": { + "node": ">= 0.6" + } + }, + "node_modules/eventsource": { + "version": "3.0.7", + "resolved": "https://registry.npmjs.org/eventsource/-/eventsource-3.0.7.tgz", + "integrity": + "sha512-CRT1WTyuQoD771GW56XEZFQ/ZoSfWid1alKGDYMmkt2yl8UXrVR4pspqWNEcqKvVIzg6PAltWjxcSSPrboA4iA==", + "license": "MIT", + "dependencies": { + "eventsource-parser": "^3.0.1" + }, + "engines": { + "node": ">=18.0.0" + } + }, + "node_modules/eventsource-parser": { + "version": "3.1.0", + "resolved": "https://registry.npmjs.org/eventsource-parser/-/eventsource-parser-3.1.0.tgz", + "integrity": + "sha512-kJezFj9YFAMLeORyi7aCLxLbD5/qWMQnoMVlVPyHIll7lgRJCc3JVln9Vgl9nwQi0YkMnhdGTMNn7CkRRAptMg==", + "license": "MIT", + "engines": { + "node": ">=18.0.0" + } + }, + "node_modules/express": { + "version": "5.2.1", + "resolved": "https://registry.npmjs.org/express/-/express-5.2.1.tgz", + "integrity": + "sha512-hIS4idWWai69NezIdRt2xFVofaF4j+6INOpJlVOLDO8zXGpUVEVzIYk12UUi2JzjEzWL3IOAxcTubgz9Po0yXw==", + "license": "MIT", + "dependencies": { + "accepts": "^2.0.0", + "body-parser": "^2.2.1", + "content-disposition": "^1.0.0", + "content-type": "^1.0.5", + "cookie": "^0.7.1", + "cookie-signature": "^1.2.1", + "debug": "^4.4.0", + "depd": "^2.0.0", + "encodeurl": "^2.0.0", + "escape-html": "^1.0.3", + "etag": "^1.8.1", + "finalhandler": "^2.1.0", + "fresh": "^2.0.0", + "http-errors": "^2.0.0", + "merge-descriptors": "^2.0.0", + "mime-types": "^3.0.0", + "on-finished": "^2.4.1", + "once": "^1.4.0", + "parseurl": "^1.3.3", + "proxy-addr": "^2.0.7", + "qs": "^6.14.0", + "range-parser": "^1.2.1", + "router": "^2.2.0", + "send": "^1.1.0", + "serve-static": "^2.2.0", + "statuses": "^2.0.1", + "type-is": "^2.0.1", + "vary": "^1.1.2" + }, + "engines": { + "node": ">= 18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/express-rate-limit": { + "version": "8.6.0", + "resolved": "https://registry.npmjs.org/express-rate-limit/-/express-rate-limit-8.6.0.tgz", + "integrity": + "sha512-XKJXDsASUOo0LLtFwW5hCcQGH0N4WQc/Rn8/Pvoia+TJFOkkFPvrtW9lZOeeNcxQJspvOIERMwiRLsVFlhHEkA==", + "license": "MIT", + "dependencies": { + "debug": "^4.4.3", + "ip-address": "^10.2.0" + }, + "engines": { + "node": ">= 16" + }, + "funding": { + "url": "https://github.com/sponsors/express-rate-limit" + }, + "peerDependencies": { + "express": ">= 4.11" + } + }, + "node_modules/fast-deep-equal": { + "version": "3.1.3", + "resolved": "https://registry.npmjs.org/fast-deep-equal/-/fast-deep-equal-3.1.3.tgz", + "integrity": + "sha512-f3qQ9oQy9j2AhBe/H9VC91wLmKBCCU/gDOnKNAYG5hswO7BLKj09Hc5HYNz9cGI++xlpDCIgDaitVs03ATR84Q==", + "license": "MIT" + }, + "node_modules/fast-uri": { + "version": "3.1.4", + "resolved": "https://registry.npmjs.org/fast-uri/-/fast-uri-3.1.4.tgz", + "integrity": + "sha512-8JnbkQ4juDyvYs4mgFGQqg4yCYtFDtUtmp2QIQq11ZZe5CFQ5wcqm1rqDgAh/QdMySuBnPzMUiJUNZG5N/AiQw==", + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/fastify" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/fastify" + } + ], + "license": "BSD-3-Clause" + }, + "node_modules/finalhandler": { + "version": "2.1.1", + "resolved": "https://registry.npmjs.org/finalhandler/-/finalhandler-2.1.1.tgz", + "integrity": + "sha512-S8KoZgRZN+a5rNwqTxlZZePjT/4cnm0ROV70LedRHZ0p8u9fRID0hJUZQpkKLzro8LfmC8sx23bY6tVNxv8pQA==", + "license": "MIT", + "dependencies": { + "debug": "^4.4.0", + "encodeurl": "^2.0.0", + "escape-html": "^1.0.3", + "on-finished": "^2.4.1", + "parseurl": "^1.3.3", + "statuses": "^2.0.1" + }, + "engines": { + "node": ">= 18.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/forwarded": { + "version": "0.2.0", + "resolved": "https://registry.npmjs.org/forwarded/-/forwarded-0.2.0.tgz", + "integrity": + "sha512-buRG0fpBtRHSTCOASe6hD258tEubFoRLb4ZNA6NxMVHNw2gOcwHo9wyablzMzOA5z9xA9L1KNjk/Nt6MT9aYow==", + "license": "MIT", + "engines": { + "node": ">= 0.6" + } + }, + "node_modules/fresh": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/fresh/-/fresh-2.0.0.tgz", + "integrity": + "sha512-Rx/WycZ60HOaqLKAi6cHRKKI7zxWbJ31MhntmtwMoaTeF7XFH9hhBp8vITaMidfljRQ6eYWCKkaTK+ykVJHP2A==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/function-bind": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/function-bind/-/function-bind-1.1.2.tgz", + "integrity": + "sha512-7XHNxH7qX9xG5mIwxkhumTox/MIRNcOgDrxWsMt2pAr23WHp6MrRlN7FBSFpCpr+oVO0F744iUgR82nJMfG2SA==", + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/get-intrinsic": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/get-intrinsic/-/get-intrinsic-1.3.0.tgz", + "integrity": + "sha512-9fSjSaos/fRIVIp+xSJlE6lfwhES7LNtKaCBIamHsjr2na1BiABJPo0mOjjz8GJDURarmCPGqaiVg5mfjb98CQ==", + "license": "MIT", + "dependencies": { + "call-bind-apply-helpers": "^1.0.2", + "es-define-property": "^1.0.1", + "es-errors": "^1.3.0", + "es-object-atoms": "^1.1.1", + "function-bind": "^1.1.2", + "get-proto": "^1.0.1", + "gopd": "^1.2.0", + "has-symbols": "^1.1.0", + "hasown": "^2.0.2", + "math-intrinsics": "^1.1.0" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/get-proto": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/get-proto/-/get-proto-1.0.1.tgz", + "integrity": + "sha512-sTSfBjoXBp89JvIKIefqw7U2CCebsc74kiY6awiGogKtoSGbgjYE/G/+l9sF3MWFPNc9IcoOC4ODfKHfxFmp0g==", + "license": "MIT", + "dependencies": { + "dunder-proto": "^1.0.1", + "es-object-atoms": "^1.0.0" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/gopd": { + "version": "1.2.0", + "resolved": "https://registry.npmjs.org/gopd/-/gopd-1.2.0.tgz", + "integrity": + "sha512-ZUKRh6/kUFoAiTAtTYPZJ3hw9wNxx+BIBOijnlG9PnrJsCcSjs1wyyD6vJpaYtgnzDrKYRSqf3OO6Rfa93xsRg==", + "license": "MIT", + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/has-symbols": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/has-symbols/-/has-symbols-1.1.0.tgz", + "integrity": + "sha512-1cDNdwJ2Jaohmb3sg4OmKaMBwuC48sYni5HUw2DvsC8LjGTLK9h+eb1X6RyuOHe4hT0ULCW68iomhjUoKUqlPQ==", + "license": "MIT", + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/hasown": { + "version": "2.0.4", + "resolved": "https://registry.npmjs.org/hasown/-/hasown-2.0.4.tgz", + "integrity": + "sha512-T2UbfbBEF32wiepXIsMlTW9+dDYC6wMh/t/vYA4tuOMKqWz/n3vr1NFSxQiyP+zk2mXsoMA/i/7qV6LKut1t1A==", + "license": "MIT", + "dependencies": { + "function-bind": "^1.1.2" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/hono": { + "version": "4.12.31", + "resolved": "https://registry.npmjs.org/hono/-/hono-4.12.31.tgz", + "integrity": + "sha512-zJIHFrl6bq3RDd2YusFNCDlM8qUprxKswyi/OPzPyzKDdyBXDqWx8bZlZ7R+saTdSTatUmb3O7K4SspGPaEOQg==", + "license": "MIT", + "engines": { + "node": ">=16.9.0" + } + }, + "node_modules/http-errors": { + "version": "2.0.1", + "resolved": "https://registry.npmjs.org/http-errors/-/http-errors-2.0.1.tgz", + "integrity": + "sha512-4FbRdAX+bSdmo4AUFuS0WNiPz8NgFt+r8ThgNWmlrjQjt1Q7ZR9+zTlce2859x4KSXrwIsaeTqDoKQmtP8pLmQ==", + "license": "MIT", + "dependencies": { + "depd": "~2.0.0", + "inherits": "~2.0.4", + "setprototypeof": "~1.2.0", + "statuses": "~2.0.2", + "toidentifier": "~1.0.1" + }, + "engines": { + "node": ">= 0.8" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/iconv-lite": { + "version": "0.7.3", + "resolved": "https://registry.npmjs.org/iconv-lite/-/iconv-lite-0.7.3.tgz", + "integrity": + "sha512-IKXpvIzjnC9XTAUbVBcMfGS0EPaIXtW6v+zr+RRp+hqULEpo0owZax6wyRwPOJbWbzjYspQwusTsfVr0ifh4uQ==", + "license": "MIT", + "dependencies": { + "safer-buffer": ">= 2.1.2 < 3.0.0" + }, + "engines": { + "node": ">=0.10.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/inherits": { + "version": "2.0.4", + "resolved": "https://registry.npmjs.org/inherits/-/inherits-2.0.4.tgz", + "integrity": + "sha512-k/vGaX4/Yla3WzyMCvTQOXYeIHvqOKtnqBduzTHpzpQZzAskKMhZ2K+EnBiSM9zGSoIFeMpXKxa4dYeZIQqewQ==", + "license": "ISC" + }, + "node_modules/ip-address": { + "version": "10.2.0", + "resolved": "https://registry.npmjs.org/ip-address/-/ip-address-10.2.0.tgz", + "integrity": + "sha512-/+S6j4E9AHvW9SWMSEY9Xfy66O5PWvVEJ08O0y5JGyEKQpojb0K0GKpz/v5HJ/G0vi3D2sjGK78119oXZeE0qA==", + "license": "MIT", + "engines": { + "node": ">= 12" + } + }, + "node_modules/ipaddr.js": { + "version": "1.9.1", + "resolved": "https://registry.npmjs.org/ipaddr.js/-/ipaddr.js-1.9.1.tgz", + "integrity": + "sha512-0KI/607xoxSToH7GjN1FfSbLoU0+btTicjsQSWQlh/hZykN8KpmMf7uYwPW3R+akZ6R/w18ZlXSHBYXiYUPO3g==", + "license": "MIT", + "engines": { + "node": ">= 0.10" + } + }, + "node_modules/is-promise": { + "version": "4.0.0", + "resolved": "https://registry.npmjs.org/is-promise/-/is-promise-4.0.0.tgz", + "integrity": + "sha512-hvpoI6korhJMnej285dSg6nu1+e6uxs7zG3BYAm5byqDsgJNWwxzM6z6iZiAgQR4TJ30JmBTOwqZUw3WlyH3AQ==", + "license": "MIT" + }, + "node_modules/isexe": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/isexe/-/isexe-2.0.0.tgz", + "integrity": + "sha512-RHxMLp9lnKHGHRng9QFhRCMbYAcVpn69smSGcq3f36xjgVVWThj4qqLbTLlq7Ssj8B+fIQ1EuCEGI2lKsyQeIw==", + "license": "ISC" + }, + "node_modules/jose": { + "version": "6.2.3", + "resolved": "https://registry.npmjs.org/jose/-/jose-6.2.3.tgz", + "integrity": + "sha512-YYVDInQKFJfR/xa3ojUTl8c2KoTwiL1R5Wg9YCydwH0x0B9grbzlg5HC7mMjCtUJjbQ/YnGEZIhI5tCgfTb4Hw==", + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/panva" + } + }, + "node_modules/json-schema-traverse": { + "version": "1.0.0", + "resolved": + "https://registry.npmjs.org/json-schema-traverse/-/json-schema-traverse-1.0.0.tgz", + "integrity": + "sha512-NM8/P9n3XjXhIZn1lLhkFaACTOURQXjWhV4BA/RnOv8xvgqtqpAX9IO4mRQxSx1Rlo4tqzeqb0sOlruaOy3dug==", + "license": "MIT" + }, + "node_modules/json-schema-typed": { + "version": "8.0.2", + "resolved": "https://registry.npmjs.org/json-schema-typed/-/json-schema-typed-8.0.2.tgz", + "integrity": + "sha512-fQhoXdcvc3V28x7C7BMs4P5+kNlgUURe2jmUT1T//oBRMDrqy1QPelJimwZGo7Hg9VPV3EQV5Bnq4hbFy2vetA==", + "license": "BSD-2-Clause" + }, + "node_modules/json-with-bigint": { + "version": "3.5.10", + "resolved": "https://registry.npmjs.org/json-with-bigint/-/json-with-bigint-3.5.10.tgz", + "integrity": + "sha512-Vcx+JVNEBts/xfcoCS69sKrOhOk/3TVlvlT+XzUOefVKnnrbYSCKpDCm10pohsJFtsJVYnwa/cXRZ4eElzaM6w==", + "license": "MIT" + }, + "node_modules/math-intrinsics": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/math-intrinsics/-/math-intrinsics-1.1.0.tgz", + "integrity": + "sha512-/IXtbwEk5HTPyEwyKX6hGkYXxM9nbj64B+ilVJnC/R6B0pH5G4V3b0pVbL7DBj4tkhBAppbQUlf6F6Xl9LHu1g==", + "license": "MIT", + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/media-typer": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/media-typer/-/media-typer-1.1.0.tgz", + "integrity": + "sha512-aisnrDP4GNe06UcKFnV5bfMNPBUw4jsLGaWwWfnH3v02GnBuXX2MCVn5RbrWo0j3pczUilYblq7fQ7Nw2t5XKw==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/merge-descriptors": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/merge-descriptors/-/merge-descriptors-2.0.0.tgz", + "integrity": + "sha512-Snk314V5ayFLhp3fkUREub6WtjBfPdCPY1Ln8/8munuLuiYhsABgBVWsozAG+MWMbVEvcdcpbi9R7ww22l9Q3g==", + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/mime-db": { + "version": "1.54.0", + "resolved": "https://registry.npmjs.org/mime-db/-/mime-db-1.54.0.tgz", + "integrity": + "sha512-aU5EJuIN2WDemCcAp2vFBfp/m4EAhWJnUNSSw0ixs7/kXbd6Pg64EmwJkNdFhB8aWt1sH2CTXrLxo/iAGV3oPQ==", + "license": "MIT", + "engines": { + "node": ">= 0.6" + } + }, + "node_modules/mime-types": { + "version": "3.0.2", + "resolved": "https://registry.npmjs.org/mime-types/-/mime-types-3.0.2.tgz", + "integrity": + "sha512-Lbgzdk0h4juoQ9fCKXW4by0UJqj+nOOrI9MJ1sSj4nI8aI2eo1qmvQEie4VD1glsS250n15LsWsYtCugiStS5A==", + "license": "MIT", + "dependencies": { + "mime-db": "^1.54.0" + }, + "engines": { + "node": ">=18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/ms": { + "version": "2.1.3", + "resolved": "https://registry.npmjs.org/ms/-/ms-2.1.3.tgz", + "integrity": + "sha512-6FlzubTLZG3J2a/NVCAleEhjzq5oxgHyaCU9yYXvcLsvoVaHJq/s5xXI6/XXP6tz7R9xAOtHnSO/tXtF3WRTlA==", + "license": "MIT" + }, + "node_modules/negotiator": { + "version": "1.0.0", + "resolved": "https://registry.npmjs.org/negotiator/-/negotiator-1.0.0.tgz", + "integrity": + "sha512-8Ofs/AUQh8MaEcrlq5xOX0CQ9ypTF5dl78mjlMNfOK08fzpgTHQRQPBxcPlEtIw0yRpws+Zo/3r+5WRby7u3Gg==", + "license": "MIT", + "engines": { + "node": ">= 0.6" + } + }, + "node_modules/object-assign": { + "version": "4.1.1", + "resolved": "https://registry.npmjs.org/object-assign/-/object-assign-4.1.1.tgz", + "integrity": + "sha512-rJgTQnkUnH1sFw8yT6VSU3zD3sWmu6sZhIseY8VX+GRu3P6F7Fu+JNDoXfklElbLJSnc3FUQHVe4cU5hj+BcUg==", + "license": "MIT", + "engines": { + "node": ">=0.10.0" + } + }, + "node_modules/object-inspect": { + "version": "1.13.4", + "resolved": "https://registry.npmjs.org/object-inspect/-/object-inspect-1.13.4.tgz", + "integrity": + "sha512-W67iLl4J2EXEGTbfeHCffrjDfitvLANg0UlX3wFUUSTx92KXRFegMHUVgSqE+wvhAbi4WqjGg9czysTV2Epbew==", + "license": "MIT", + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/on-finished": { + "version": "2.4.1", + "resolved": "https://registry.npmjs.org/on-finished/-/on-finished-2.4.1.tgz", + "integrity": + "sha512-oVlzkg3ENAhCk2zdv7IJwd/QUD4z2RxRwpkcGY8psCVcCYZNq4wYnVWALHM+brtuJjePWiYF/ClmuDr8Ch5+kg==", + "license": "MIT", + "dependencies": { + "ee-first": "1.1.1" + }, + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/once": { + "version": "1.4.0", + "resolved": "https://registry.npmjs.org/once/-/once-1.4.0.tgz", + "integrity": + "sha512-lNaJgI+2Q5URQBkccEKHTQOPaXdUxnZZElQTZY0MFUAuaEqe1E+Nyvgdz/aIyNi6Z9MzO5dv1H8n58/GELp3+w==", + "license": "ISC", + "dependencies": { + "wrappy": "1" + } + }, + "node_modules/parseurl": { + "version": "1.3.3", + "resolved": "https://registry.npmjs.org/parseurl/-/parseurl-1.3.3.tgz", + "integrity": + "sha512-CiyeOxFT/JZyN5m0z9PfXw4SCBJ6Sygz1Dpl0wqjlhDEGGBP1GnsUVEL0p63hoG1fcj3fHynXi9NYO4nWOL+qQ==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/path-key": { + "version": "3.1.1", + "resolved": "https://registry.npmjs.org/path-key/-/path-key-3.1.1.tgz", + "integrity": + "sha512-ojmeN0qd+y0jszEtoY48r0Peq5dwMEkIlCOu6Q5f41lfkswXuKtYrhgoTpLnyIcHm24Uhqx+5Tqm2InSwLhE6Q==", + "license": "MIT", + "engines": { + "node": ">=8" + } + }, + "node_modules/path-to-regexp": { + "version": "8.4.2", + "resolved": "https://registry.npmjs.org/path-to-regexp/-/path-to-regexp-8.4.2.tgz", + "integrity": + "sha512-qRcuIdP69NPm4qbACK+aDogI5CBDMi1jKe0ry5rSQJz8JVLsC7jV8XpiJjGRLLol3N+R5ihGYcrPLTno6pAdBA==", + "license": "MIT", + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/pkce-challenge": { + "version": "5.0.1", + "resolved": "https://registry.npmjs.org/pkce-challenge/-/pkce-challenge-5.0.1.tgz", + "integrity": + "sha512-wQ0b/W4Fr01qtpHlqSqspcj3EhBvimsdh0KlHhH8HRZnMsEa0ea2fTULOXOS9ccQr3om+GcGRk4e+isrZWV8qQ==", + "license": "MIT", + "engines": { + "node": ">=16.20.0" + } + }, + "node_modules/proxy-addr": { + "version": "2.0.7", + "resolved": "https://registry.npmjs.org/proxy-addr/-/proxy-addr-2.0.7.tgz", + "integrity": + "sha512-llQsMLSUDUPT44jdrU/O37qlnifitDP+ZwrmmZcoSKyLKvtZxpyV0n2/bD/N4tBAAZ/gJEdZU7KMraoK1+XYAg==", + "license": "MIT", + "dependencies": { + "forwarded": "0.2.0", + "ipaddr.js": "1.9.1" + }, + "engines": { + "node": ">= 0.10" + } + }, + "node_modules/qs": { + "version": "6.15.3", + "resolved": "https://registry.npmjs.org/qs/-/qs-6.15.3.tgz", + "integrity": + "sha512-O9gl3zCl5h5blw1KGUzQKhA5oUXSl8rwUIM5o0S3nCXMliSvy5Dzx7/DJcI+SwgICv+IneSZwhBh1oSyEHA71A==", + "license": "BSD-3-Clause", + "dependencies": { + "es-define-property": "^1.0.1", + "side-channel": "^1.1.1" + }, + "engines": { + "node": ">=0.6" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/range-parser": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/range-parser/-/range-parser-1.3.0.tgz", + "integrity": + "sha512-hek2mFQpPuI4E1BBKrSto+BU3e3x4xuarsbiwr3+lf7p44juvFMV0XFWQAP3xUyqXA4RrXLIoaSUGbSt056ZMw==", + "license": "MIT", + "engines": { + "node": ">= 0.6" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/raw-body": { + "version": "3.0.2", + "resolved": "https://registry.npmjs.org/raw-body/-/raw-body-3.0.2.tgz", + "integrity": + "sha512-K5zQjDllxWkf7Z5xJdV0/B0WTNqx6vxG70zJE4N0kBs4LovmEYWJzQGxC9bS9RAKu3bgM40lrd5zoLJ12MQ5BA==", + "license": "MIT", + "dependencies": { + "bytes": "~3.1.2", + "http-errors": "~2.0.1", + "iconv-lite": "~0.7.0", + "unpipe": "~1.0.0" + }, + "engines": { + "node": ">= 0.10" + } + }, + "node_modules/require-from-string": { + "version": "2.0.2", + "resolved": + "https://registry.npmjs.org/require-from-string/-/require-from-string-2.0.2.tgz", + "integrity": + "sha512-Xf0nWe6RseziFMu+Ap9biiUbmplq6S9/p+7w7YXP/JBHhrUDDUhwa+vANyubuqfZWTveU//DYVGsDG7RKL/vEw==", + "license": "MIT", + "engines": { + "node": ">=0.10.0" + } + }, + "node_modules/router": { + "version": "2.2.0", + "resolved": "https://registry.npmjs.org/router/-/router-2.2.0.tgz", + "integrity": + "sha512-nLTrUKm2UyiL7rlhapu/Zl45FwNgkZGaCpZbIHajDYgwlJCOzLSk+cIPAnsEqV955GjILJnKbdQC1nVPz+gAYQ==", + "license": "MIT", + "dependencies": { + "debug": "^4.4.0", + "depd": "^2.0.0", + "is-promise": "^4.0.0", + "parseurl": "^1.3.3", + "path-to-regexp": "^8.0.0" + }, + "engines": { + "node": ">= 18" + } + }, + "node_modules/safer-buffer": { + "version": "2.1.2", + "resolved": "https://registry.npmjs.org/safer-buffer/-/safer-buffer-2.1.2.tgz", + "integrity": + "sha512-YZo3K82SD7Riyi0E1EQPojLz7kpepnSQI9IyPbHHg1XXXevb5dJI7tpyN2ADxGcQbHG7vcyRHk0cbwqcQriUtg==", + "license": "MIT" + }, + "node_modules/send": { + "version": "1.2.1", + "resolved": "https://registry.npmjs.org/send/-/send-1.2.1.tgz", + "integrity": + "sha512-1gnZf7DFcoIcajTjTwjwuDjzuz4PPcY2StKPlsGAQ1+YH20IRVrBaXSWmdjowTJ6u8Rc01PoYOGHXfP1mYcZNQ==", + "license": "MIT", + "dependencies": { + "debug": "^4.4.3", + "encodeurl": "^2.0.0", + "escape-html": "^1.0.3", + "etag": "^1.8.1", + "fresh": "^2.0.0", + "http-errors": "^2.0.1", + "mime-types": "^3.0.2", + "ms": "^2.1.3", + "on-finished": "^2.4.1", + "range-parser": "^1.2.1", + "statuses": "^2.0.2" + }, + "engines": { + "node": ">= 18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/serve-static": { + "version": "2.2.1", + "resolved": "https://registry.npmjs.org/serve-static/-/serve-static-2.2.1.tgz", + "integrity": + "sha512-xRXBn0pPqQTVQiC8wyQrKs2MOlX24zQ0POGaj0kultvoOCstBQM5yvOhAVSUwOMjQtTvsPWoNCHfPGwaaQJhTw==", + "license": "MIT", + "dependencies": { + "encodeurl": "^2.0.0", + "escape-html": "^1.0.3", + "parseurl": "^1.3.3", + "send": "^1.2.0" + }, + "engines": { + "node": ">= 18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/setprototypeof": { + "version": "1.2.0", + "resolved": "https://registry.npmjs.org/setprototypeof/-/setprototypeof-1.2.0.tgz", + "integrity": + "sha512-E5LDX7Wrp85Kil5bhZv46j8jOeboKq5JMmYM3gVGdGH8xFpPWXUMsNrlODCrkoxMEeNi/XZIwuRvY4XNwYMJpw==", + "license": "ISC" + }, + "node_modules/shebang-command": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/shebang-command/-/shebang-command-2.0.0.tgz", + "integrity": + "sha512-kHxr2zZpYtdmrN1qDjrrX/Z1rR1kG8Dx+gkpK1G4eXmvXswmcE1hTWBWYUzlraYw1/yZp6YuDY77YtvbN0dmDA==", + "license": "MIT", + "dependencies": { + "shebang-regex": "^3.0.0" + }, + "engines": { + "node": ">=8" + } + }, + "node_modules/shebang-regex": { + "version": "3.0.0", + "resolved": "https://registry.npmjs.org/shebang-regex/-/shebang-regex-3.0.0.tgz", + "integrity": + "sha512-7++dFhtcx3353uBaq8DDR4NuxBetBzC7ZQOhmTQInHEd6bSrXdiEyzCvG07Z44UYdLShWUyXt5M/yhz8ekcb1A==", + "license": "MIT", + "engines": { + "node": ">=8" + } + }, + "node_modules/side-channel": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/side-channel/-/side-channel-1.1.1.tgz", + "integrity": + "sha512-6x6dK6zJdpTzF4sQeNYxwtvBzf6Eg4GtlesS94HOvTudUeyK2WXAaIfmDgsyslYrRBeFIlsi54AYsFGUuhmvrQ==", + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0", + "object-inspect": "^1.13.4", + "side-channel-list": "^1.0.1", + "side-channel-map": "^1.0.1", + "side-channel-weakmap": "^1.0.2" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/side-channel-list": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/side-channel-list/-/side-channel-list-1.0.1.tgz", + "integrity": + "sha512-mjn/0bi/oUURjc5Xl7IaWi/OJJJumuoJFQJfDDyO46+hBWsfaVM65TBHq2eoZBhzl9EchxOijpkbRC8SVBQU0w==", + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0", + "object-inspect": "^1.13.4" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/side-channel-map": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/side-channel-map/-/side-channel-map-1.0.1.tgz", + "integrity": + "sha512-VCjCNfgMsby3tTdo02nbjtM/ewra6jPHmpThenkTYh8pG9ucZ/1P8So4u4FGBek/BjpOVsDCMoLA/iuBKIFXRA==", + "license": "MIT", + "dependencies": { + "call-bound": "^1.0.2", + "es-errors": "^1.3.0", + "get-intrinsic": "^1.2.5", + "object-inspect": "^1.13.3" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/side-channel-weakmap": { + "version": "1.0.2", + "resolved": + "https://registry.npmjs.org/side-channel-weakmap/-/side-channel-weakmap-1.0.2.tgz", + "integrity": + "sha512-WPS/HvHQTYnHisLo9McqBHOJk2FkHO/tlpvldyrnem4aeQp4hai3gythswg6p01oSoTl58rcpiFAjF2br2Ak2A==", + "license": "MIT", + "dependencies": { + "call-bound": "^1.0.2", + "es-errors": "^1.3.0", + "get-intrinsic": "^1.2.5", + "object-inspect": "^1.13.3", + "side-channel-map": "^1.0.1" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/statuses": { + "version": "2.0.2", + "resolved": "https://registry.npmjs.org/statuses/-/statuses-2.0.2.tgz", + "integrity": + "sha512-DvEy55V3DB7uknRo+4iOGT5fP1slR8wQohVdknigZPMpMstaKJQWhwiYBACJE3Ul2pTnATihhBYnRhZQHGBiRw==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/toidentifier": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/toidentifier/-/toidentifier-1.0.1.tgz", + "integrity": + "sha512-o5sSPKEkg/DIQNmH43V0/uerLrpzVedkUh8tGNvaeXpfpuwjKenlSox/2O/BTlZUtEe+JG7s5YhEz608PlAHRA==", + "license": "MIT", + "engines": { + "node": ">=0.6" + } + }, + "node_modules/type-is": { + "version": "2.1.0", + "resolved": "https://registry.npmjs.org/type-is/-/type-is-2.1.0.tgz", + "integrity": + "sha512-faYHw0anBbc/kWF3zFTEnxSFOAGUX9GFbOBthvDdLsIlEoWOFOtS0zgCiQYwIskL9iGXZL3kAXD8OoZ4GmMATA==", + "license": "MIT", + "dependencies": { + "content-type": "^2.0.0", + "media-typer": "^1.1.0", + "mime-types": "^3.0.0" + }, + "engines": { + "node": ">= 18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/type-is/node_modules/content-type": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/content-type/-/content-type-2.0.0.tgz", + "integrity": + "sha512-j/O/d7GcZCyNl7/hwZAb606rzqkyvaDctLmckbxLzHvFBzTJHuGEdodATcP3yIRoDrLHkIATJuvzbFlp/ki2cQ==", + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/express" + } + }, + "node_modules/undici": { + "version": "7.28.0", + "resolved": "https://registry.npmjs.org/undici/-/undici-7.28.0.tgz", + "integrity": + "sha512-cRZYrTDwWznlnRiPjggAGxZXanty6M8RV1ff8Wm4LWXBp7/IG8v5DnOm74DtUBp9OONpK75YlPnIjQqX0dBDtA==", + "license": "MIT", + "engines": { + "node": ">=20.18.1" + } + }, + "node_modules/universal-user-agent": { + "version": "7.0.3", + "resolved": + "https://registry.npmjs.org/universal-user-agent/-/universal-user-agent-7.0.3.tgz", + "integrity": + "sha512-TmnEAEAsBJVZM/AADELsK76llnwcf9vMKuPz8JflO1frO8Lchitr0fNaN9d+Ap0BjKtqWqd/J17qeDnXh8CL2A==", + "license": "ISC" + }, + "node_modules/unpipe": { + "version": "1.0.0", + "resolved": "https://registry.npmjs.org/unpipe/-/unpipe-1.0.0.tgz", + "integrity": + "sha512-pjy2bYhSsufwWlKwPc+l3cN7+wuJlK6uz0YdJEOlQDbl6jo/YlPi4mb8agUkVC8BF7V8NuzeyPNqRksA3hztKQ==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/vary": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/vary/-/vary-1.1.2.tgz", + "integrity": + "sha512-BNGbWLfd0eUPabhkXUVm0j8uuvREyTh5ovRa/dyow/BqAbZJyC+5fU+IzQOzmAKzYqYRAISoRhdQr3eIZ/PXqg==", + "license": "MIT", + "engines": { + "node": ">= 0.8" + } + }, + "node_modules/which": { + "version": "2.0.2", + "resolved": "https://registry.npmjs.org/which/-/which-2.0.2.tgz", + "integrity": + "sha512-BLI3Tl1TW3Pvl70l3yq3Y64i+awpwXqsGBYWkkqMtnbXgrMD+yj7rhW0kuEDxzJaYXGjEW5ogapKNMEKNMjibA==", + "license": "ISC", + "dependencies": { + "isexe": "^2.0.0" + }, + "bin": { + "node-which": "bin/node-which" + }, + "engines": { + "node": ">= 8" + } + }, + "node_modules/wrappy": { + "version": "1.0.2", + "resolved": "https://registry.npmjs.org/wrappy/-/wrappy-1.0.2.tgz", + "integrity": + "sha512-l4Sp/DRseor9wL6EvV2+TuQn63dMkPjZ/sp9XkghTEbV9KlPS1xUsZ3u7/IQO4wxtcFB4bgpQPRcR3QCvezPcQ==", + "license": "ISC" + }, + "node_modules/yaml": { + "version": "2.9.0", + "resolved": "https://registry.npmjs.org/yaml/-/yaml-2.9.0.tgz", + "integrity": + "sha512-2AvhNX3mb8zd6Zy7INTtSpl1F15HW6Wnqj0srWlkKLcpYl/gMIMJiyuGq2KeI2YFxUPjdlB+3Lc10seMLtL4cA==", + "license": "ISC", + "bin": { + "yaml": "bin.mjs" + }, + "engines": { + "node": ">= 14.6" + }, + "funding": { + "url": "https://github.com/sponsors/eemeli" + } + }, + "node_modules/zod": { + "version": "4.4.3", + "resolved": "https://registry.npmjs.org/zod/-/zod-4.4.3.tgz", + "integrity": + "sha512-ytENFjIJFl2UwYglde2jchW2Hwm4GJFLDiSXWdTrJQBIN9Fcyp7n4DhxJEiWNAJMV1/BqWfW/kkg71UDcHJyTQ==", + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/colinhacks" + } + }, + "node_modules/zod-to-json-schema": { + "version": "3.25.2", + "resolved": "https://registry.npmjs.org/zod-to-json-schema/-/zod-to-json-schema-3.25.2.tgz", + "integrity": + "sha512-O/PgfnpT1xKSDeQYSCfRI5Gy3hPf91mKVDuYLUHZJMiDFptvP41MSnWofm8dnCm0256ZNfZIM7DSzuSMAFnjHA==", + "license": "ISC", + "peerDependencies": { + "zod": "^3.25.28 || ^4" + } + } + } +} diff --git a/conformance/runner/package.json b/conformance/runner/package.json new file mode 100644 index 0000000..8b24c60 --- /dev/null +++ b/conformance/runner/package.json @@ -0,0 +1,11 @@ +{ + "name": "mcp-cpp-sdk-conformance-runner", + "private": true, + "version": "0.0.0", + "engines": { + "node": ">=20" + }, + "dependencies": { + "@modelcontextprotocol/conformance": "0.1.16" + } +} diff --git a/conformance/summarize.py b/conformance/summarize.py new file mode 100644 index 0000000..4e9cb6b --- /dev/null +++ b/conformance/summarize.py @@ -0,0 +1,174 @@ +#!/usr/bin/env python3 +"""Create scenario-level conformance evidence from runner check files.""" + +from __future__ import annotations + +import argparse +import json +import re +from pathlib import Path + + +TIMESTAMP_SUFFIX = re.compile(r"-\d{4}-\d{2}-\d{2}T\d{2}-\d{2}-\d{2}-\d{3}Z$") + +EXPECTED_SCENARIOS = { + "server": { + "completion-complete", + "dns-rebinding-protection", + "elicitation-sep1034-defaults", + "elicitation-sep1330-enums", + "logging-set-level", + "ping", + "prompts-get-embedded-resource", + "prompts-get-simple", + "prompts-get-with-args", + "prompts-get-with-image", + "prompts-list", + "resources-list", + "resources-read-binary", + "resources-read-text", + "resources-subscribe", + "resources-templates-read", + "resources-unsubscribe", + "server-initialize", + "server-sse-multiple-streams", + "tools-call-audio", + "tools-call-elicitation", + "tools-call-embedded-resource", + "tools-call-error", + "tools-call-image", + "tools-call-mixed-content", + "tools-call-sampling", + "tools-call-simple-text", + "tools-call-with-logging", + "tools-call-with-progress", + "tools-list", + }, + "client": { + "auth/basic-cimd", + "auth/metadata-default", + "auth/metadata-var1", + "auth/metadata-var2", + "auth/metadata-var3", + "auth/pre-registration", + "auth/scope-from-scopes-supported", + "auth/scope-from-www-authenticate", + "auth/scope-omitted-when-undefined", + "auth/scope-retry-limit", + "auth/scope-step-up", + "auth/token-endpoint-auth-basic", + "auth/token-endpoint-auth-none", + "auth/token-endpoint-auth-post", + "elicitation-sep1034-client-defaults", + "initialize", + "sse-retry", + "tools_call", + }, +} + + +def scenario_name(suite: str, root: Path, checks_file: Path) -> str: + name = checks_file.parent.relative_to(root).as_posix() + name = TIMESTAMP_SUFFIX.sub("", name) + if suite == "server" and name.startswith("server-"): + name = name.removeprefix("server-") + return name + + +def summarize_suite(suite: str, root: Path) -> dict[str, object]: + emitted: dict[str, list[Path]] = {} + for checks_file in sorted(root.rglob("checks.json")): + name = scenario_name(suite, root, checks_file) + emitted.setdefault(name, []).append(checks_file) + + scenarios: list[dict[str, object]] = [] + scenario_names = EXPECTED_SCENARIOS[suite] | emitted.keys() + for name in sorted(scenario_names): + checks_files = emitted.get(name, []) + statuses: list[str] = [] + for checks_file in checks_files: + checks = json.loads(checks_file.read_text(encoding="utf-8")) + statuses.extend(check.get("status", "UNKNOWN") for check in checks) + if not checks_files: + statuses = ["MISSING"] + elif len(checks_files) > 1: + statuses.append("DUPLICATE") + + passed = ( + name in EXPECTED_SCENARIOS[suite] + and len(checks_files) == 1 + and "FAILURE" not in statuses + and "SUCCESS" in statuses + ) + scenarios.append( + { + "name": name, + "passed": passed, + "statuses": statuses, + "checks_files": [str(path.relative_to(root)) for path in checks_files], + } + ) + + passed_count = sum(bool(scenario["passed"]) for scenario in scenarios) + total = len(scenarios) + return { + "passed": passed_count, + "total": total, + "pass_rate_percent": round(100.0 * passed_count / total, 1) if total else 0.0, + "failed_scenarios": [ + scenario["name"] for scenario in scenarios if not scenario["passed"] + ], + "scenarios": scenarios, + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--server-results", type=Path, required=True) + parser.add_argument("--client-results", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--server-status", type=int, required=True) + parser.add_argument("--client-status", type=int, required=True) + args = parser.parse_args() + + suites = { + "server": summarize_suite("server", args.server_results), + "client": summarize_suite("client", args.client_results), + } + summary = { + "runner": "@modelcontextprotocol/conformance@0.1.16", + "protocol_version": "2025-11-25", + "runner_exit_status": { + "server": args.server_status, + "client": args.client_status, + }, + "suites": suites, + } + + args.output_dir.mkdir(parents=True, exist_ok=True) + (args.output_dir / "summary.json").write_text( + json.dumps(summary, indent=2) + "\n", encoding="utf-8" + ) + + markdown = [ + "# MCP conformance regression baseline", + "", + "A green runner status means the checked-in expected-failure baseline did not drift;", + "it does not mean that an MCP SDK tier has been achieved. The pass rate follows the", + "tier-check convention: warnings do not fail a scenario, while missing evidence does.", + "", + "| Suite | Passed scenarios | Total scenarios | Pass rate | Runner status |", + "|---|---:|---:|---:|---:|", + ] + for suite in ("server", "client"): + result = suites[suite] + markdown.append( + f"| {suite} | {result['passed']} | {result['total']} | " + f"{result['pass_rate_percent']}% | {summary['runner_exit_status'][suite]} |" + ) + markdown.append("") + (args.output_dir / "summary.md").write_text("\n".join(markdown), encoding="utf-8") + + +if __name__ == "__main__": + main() diff --git a/docs/api/index.rst b/docs/api/index.rst index db7179c..d9a6287 100644 --- a/docs/api/index.rst +++ b/docs/api/index.rst @@ -105,9 +105,10 @@ Task Protocol Types -------------- -The SDK includes comprehensive protocol and supporting types for all MCP -messages. They are available through the complete generated index below, -which avoids duplicating the class-level API sections above. +The SDK includes typed protocol and supporting types for its implemented MCP +surface. They are available through the generated index below, which avoids +duplicating the class-level API sections above. The pinned conformance results, +not this index, are the source of truth for current protocol coverage. RequestId ^^^^^^^^^ diff --git a/docs/architecture.rst b/docs/architecture.rst index 13225d8..a2083fe 100644 --- a/docs/architecture.rst +++ b/docs/architecture.rst @@ -9,20 +9,24 @@ Design Principles The mcp-cpp-sdk is built on several core design principles: -Compiled Static Library -^^^^^^^^^^^^^^^^^^^^^^^^ +Compiled library boundaries +^^^^^^^^^^^^^^^^^^^^^^^^^^^ -The SDK is a compiled static library (``libmcp-cpp-sdk.a``). Implementation -details — including all Boost.Asio and Boost.Beast usage — are hidden behind -PImpl boundaries in ``.cpp`` files. Only ``core.hpp`` retains a direct Boost -include, because the ``Task`` template alias must be visible at call sites. +The SDK builds both shared and static library variants. Much of the client, +server, and network-transport runtime is implemented in ``.cpp`` files, while +serialization helpers and typed handler templates remain in public headers. +The transport and OAuth runtime implementations live in compiled translation +units. The async and transport APIs still expose Boost.Asio types and +Boost.Beast aliases, so consumers parse the corresponding Boost headers. This design: -* Reduces consumer compile times — Boost headers are not parsed for every - translation unit that includes an SDK header -* Isolates ABI-unstable Boost internals behind a stable C++ interface -* Keeps public headers free of Boost types except ``Task`` +* Keeps implementation state behind PImpl where it materially reduces public + surface area +* Places non-template runtime behavior in compiled translation units where the + public API does not require an inline definition +* Preserves templates in headers where C++ requires their definitions at the + point of instantiation RAII and Value Semantics ^^^^^^^^^^^^^^^^^^^^^^^^^ @@ -30,9 +34,13 @@ RAII and Value Semantics The SDK follows modern C++ best practices: * **RAII**: Resources (sockets, timers) are owned by objects and cleaned up automatically -* **Move semantics**: Expensive objects like ``Server`` and ``Client`` are move-only -* **Smart pointers**: ``std::unique_ptr`` for ownership, ``std::shared_ptr`` where needed -* **No raw pointers**: All memory is managed automatically +* **Stable ownership**: ``Server`` and ``Client`` are non-copyable and + non-movable; callers keep them alive while their ``run()`` or ``connect()`` + tasks are active +* **Smart pointers**: implementation, transport, and request state that must + survive suspension is shared with the relevant coroutines +* **Explicit lifetime contracts**: request-scoped non-owning views are paired + with an owning session or runtime state Coroutine-Based Async ^^^^^^^^^^^^^^^^^^^^^^ @@ -52,7 +60,9 @@ The SDK emphasizes compile-time safety: * **Concepts**: ``JsonSerializable`` concept ensures types are JSON-compatible * **Strong typing**: Protocol types are structs, not raw JSON * **Template metaprogramming**: Handler signatures validated at compile time -* **No ``void*`` or type erasure leaks**: Type erasure is internal only +* **Typed registration helpers**: Templates validate handler signatures, then + bridge to the public ``TypeErasedHandler`` middleware interface and compiled + runtime Component Overview ------------------ @@ -68,14 +78,15 @@ Server * Handles JSON-RPC message dispatch * Executes handler functions with proper context * Supports multiple handler signatures (sync/async, with/without context) -* Thread-safe via Boost.Asio strand +* Serializes each server session through a Boost.Asio strand **Key responsibilities**: * Protocol compliance (MCP handshake, request handling) * Handler type erasure and invocation * ``ensure_async_handler`` — automatically normalizes handler signatures (sync/async, with/without ``Context``) into unified async form; called internally by ``add_tool``, ``add_resource``, etc. -* Error handling: exceptions thrown during request dispatch are automatically caught and returned as ``g_INTERNAL_ERROR`` (-32603) JSON-RPC error responses +* Error handling: tool-handler exceptions become ``CallToolResult`` values with + ``isError=true``; other request-dispatch failures use JSON-RPC error responses * Context creation and lifecycle management Client @@ -87,14 +98,16 @@ Client * Sends requests and matches responses via request ID * Manages pending requests with timeout support * Handles server notifications -* Thread-safe via Boost.Asio strand +* Serializes client request, response, timeout, and close state through a + Boost.Asio strand-backed runtime **Key responsibilities**: * Request/response correlation (JSON-RPC id matching) via ``RequestId`` — a type-safe wrapper for JSON-RPC request identifiers (string or integer) * Reverse RPC: ``dispatch_incoming_request`` handles server-to-client JSON-RPC requests, enabling servers to request client capabilities (elicitation, sampling, roots) * Timeout management for requests -* Connection lifecycle (connect, run, close) +* Connection lifecycle (``connect()`` and ``close()``; the read loop is managed + internally) * Error propagation from server responses Transport Abstraction @@ -148,21 +161,22 @@ Context is passed to handlers that declare a ``Context&`` parameter. * Structured logging with severity levels * Bidirectional communication (reverse RPC) -* Request metadata (future extension point) +* Cancellation state and progress support when a request supplies a progress + token Core Types ^^^^^^^^^^ -``mcp::core`` provides fundamental types: +The core headers expose fundamental types in the ``mcp`` namespace: * ``Task`` - Coroutine return type (alias for ``boost::asio::awaitable``) -* ``LogLevel`` - Enum for log severity -* ``McpError`` - Exception type for MCP-specific errors +* ``LoggingLevel`` - Protocol enum for log severity +* ``Error`` and JSON-RPC error responses - Structured peer-visible failures Protocol Types ^^^^^^^^^^^^^^ -``mcp::protocol`` defines all MCP protocol messages: +The protocol headers expose MCP message types in the ``mcp`` namespace: * Request types: ``InitializeRequest``, ``ListToolsRequest``, ``CallToolRequest``, etc. * Response types: ``InitializeResult``, ``ListToolsResult``, ``CallToolResult``, etc. @@ -204,15 +218,13 @@ Design Decisions Why Boost.Asio? ^^^^^^^^^^^^^^^ -Boost.Asio is the de facto standard for async I/O in C++: +The SDK uses Boost.Asio for async I/O because it provides the primitives used by +the public coroutine model: -* Mature, stable, and widely used -* Excellent coroutine support (awaitable) +* C++20 coroutine support through ``awaitable`` and ``co_spawn`` * Cross-platform (Windows, Linux, macOS) -* Provides all necessary primitives (timers, streams, executors) -* Well-documented and performant - -Alternative considered: ``std::execution`` - not yet standardized or widely available. +* Timers, streams, executors, and strands under one execution model +* Established documentation and ecosystem Authentication and Authorization ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ @@ -223,7 +235,8 @@ The SDK includes OAuth 2.1 helpers in ``mcp::auth``: * ``OAuthAuthenticator``: concrete OAuth 2.0 implementation of the ``Authenticator`` interface * ``OAuthHttpClient`` for token and metadata HTTP calls * ``OAuthDiscoveryClient`` for protected resource and authorization-server discovery -* ``OAuthClientTransport`` for injecting Bearer tokens into outgoing MCP requests, with automatic token refresh on ``-32001`` or ``-32000`` error responses +* ``OAuthClientTransport`` for sending HTTP Bearer credentials and attempting + one refresh after HTTP 401 or the legacy ``g_UNAUTHORIZED`` JSON-RPC path * ``InMemoryTokenStore`` as a simple token persistence implementation This keeps authentication concerns out of the core client/server types while @@ -232,32 +245,30 @@ still allowing authenticated transports and middleware-based validation. Why nlohmann_json? ^^^^^^^^^^^^^^^^^^ -nlohmann_json is the most popular C++ JSON library: - -* Intuitive API (``j["key"] = value``) -* Excellent error messages -* Automatic type conversion -* Wide adoption in the C++ community -* Active maintenance and good performance +nlohmann_json was selected for its direct mapping between JSON and C++ protocol +types: -Alternative considered: RapidJSON - faster but more complex API. +* Object-style API (``j["key"] = value``) +* Automatic conversion through ``to_json`` and ``from_json`` +* Support for the variant and optional fields used by MCP messages -Why Compiled Static Library? -^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ +Why compiled library variants? +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ -Compiling the SDK as a static library reduces consumer build times: +The compiled variants provide implementation and linkage boundaries while +retaining the template-based public API: -* **Faster builds**: Boost headers are compiled once into the library, not - re-parsed in every consumer translation unit -* **Encapsulation**: Boost.Asio and Boost.Beast internals stay behind PImpl - boundaries and do not leak into consumer headers -* **Stable interface**: Public headers expose only standard C++ and - ``Task`` (the one unavoidable Boost type) -* **Unchanged link model**: Consumers still link a single ``mcp-cpp-sdk`` - target — no change to CMake integration +* **Encapsulation**: Internal request state and transport machinery stay behind + implementation boundaries +* **Choice of linkage**: Consumers select ``mcp::sdk_shared`` or + ``mcp::sdk_static`` explicitly, or use ``mcp::sdk`` for the configured + default +* **Template ergonomics**: Typed handlers remain available without a separate + code-generation step -Tradeoff: Cross-boundary inlining is no longer possible for PImpl'd types, -which is acceptable given that the hot paths are I/O-bound. +Tradeoff: the coroutine and executor types remain part of the source-level API, +and PImpl calls cannot be inlined across the library boundary without link-time +optimization. Why Multiple Handler Signatures? ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ @@ -274,29 +285,36 @@ Type erasure unifies these at runtime, so ``Server`` internals stay simple. Thread Safety Model ------------------- -All async operations execute on a Boost.Asio **strand**: +Client runtime transitions and each server session are serialized through +Boost.Asio strands. Separate HTTP sessions may run concurrently, and a +``StreamableHttpSessionManager`` can place tool work on a separate executor. +Application state shared by handlers therefore still needs its own +synchronization. -* Strand serializes handler execution -* No explicit locks (``std::mutex``) needed -* Simple mental model: handlers never run concurrently -* Safe to modify ``Server``/``Client`` state from handlers - -External synchronization is required only if accessing SDK objects from outside -the strand (e.g., from a different thread). +Complete server registrations before starting a session. Configure Origin and +bearer-token setters directly on ``HttpServerTransport`` or +``StreamableHttpSessionManager`` before ``listen()`` or ``run()``; +``Server::run_http()`` does not expose those transport settings. The token store +is independently synchronized, but that does not make every OAuth wrapper +operation safe for arbitrary concurrent calls. Performance Characteristics --------------------------- -* **Zero-copy**: Messages are moved, not copied (``std::string`` via move semantics) -* **Lazy serialization**: JSON parsed only when fields are accessed -* **Minimal allocations**: Handler type erasure uses ``std::function`` (one allocation) -* **No virtual dispatch in hot path**: Templates resolve at compile time +Protocol messages are parsed with ``nlohmann_json`` and cross the transport +interface as serialized strings. Typed handler dispatch is normalized once at +registration; transport dispatch remains virtual by design. Allocation and +copy behavior depends on message size, JSON values, and the selected transport, +so the project does not promise zero-copy operation. For high-throughput scenarios, consider: -* Reusing ``io_context`` across multiple connections -* Tuning ``nlohmann_json`` allocator (custom allocator support) -* Batching small messages if protocol allows +* Reusing an ``io_context`` and a multi-session HTTP manager +* Moving blocking application work to a bounded worker executor +* Measuring the complete workload with the reproducible ``benchmark/`` suite + +Benchmark results are end-to-end measurements for a pinned workload and build +configuration; they are not a universal SDK throughput guarantee. Extensibility ------------- @@ -305,7 +323,8 @@ The SDK is designed for extension: * **Custom transports**: Implement ``ITransport`` (e.g., for HTTP/2) * **Custom serialization**: Provide ``to_json``/``from_json`` for your types -* **Custom executors**: Pass any ``mcp::Runtime``-compatible executor +* **Asio executors**: Pass an executor accepted by the relevant + ``boost::asio::any_io_executor`` API * **Middleware**: Intercept handler execution for auth, logging, and request shaping Future Work @@ -314,8 +333,9 @@ Future Work Potential future enhancements: * **Connection pooling**: Reuse transports across multiple requests -* **HTTP/2 transport**: For high-performance network scenarios -* **std::execution support**: When standardized and available +* **HTTP/2 transport**: For deployments that require HTTP/2 +* **Sender/receiver interoperability**: Integration with standard execution + APIs as supported toolchains make it practical * **Batched operations**: Protocol extension for bulk requests See the GitHub issues for planned features and contributions. diff --git a/docs/concepts/middleware.rst b/docs/concepts/middleware.rst index c3992b8..580e9d6 100644 --- a/docs/concepts/middleware.rst +++ b/docs/concepts/middleware.rst @@ -89,20 +89,25 @@ Common Examples Authentication ^^^^^^^^^^^^^^ -You can create a middleware that checks for a specific header or parameter. For a full OAuth implementation, see the :doc:`oauth` concept page. +You can create middleware that checks a specific parameter. For the current +OAuth and HTTP bearer-authentication building blocks, limits, and recommended +security boundary, see :doc:`oauth`. .. code-block:: cpp server.use([](Context& ctx, const nlohmann::json& params, TypeErasedHandler next) -> Task { if (!is_authorized(params)) { - mcp::CallToolResult err; - err.isError = true; - err.content.push_back(mcp::TextContent{"Unauthorized access"}); - co_return nlohmann::json(err); + throw std::runtime_error("Unauthorized access"); } co_return co_await next(ctx, params); }); +Tool middleware returns the same application value as the wrapped handler. To +short-circuit with a tool error, throw an exception; the server converts it to +a top-level ``CallToolResult`` with ``isError=true``. Complete protocol result +objects belong in a typed ``CallToolResult`` handler or ``add_raw_tool()``, not +in an ordinary JSON middleware return value. + Logging ^^^^^^^ diff --git a/docs/concepts/oauth.rst b/docs/concepts/oauth.rst index f76be0c..aae0659 100644 --- a/docs/concepts/oauth.rst +++ b/docs/concepts/oauth.rst @@ -1,119 +1,418 @@ OAuth ===== -OAuth 2.0 and 2.1 support in the MCP C++ SDK enables secure, delegated authentication for both clients and servers. This implementation provides comprehensive tools for discovery, token management, and automatic transport-level authentication. +:cpp:class:`mcp::auth::OAuthAuthorizationManager` is the entry point for OAuth +in this SDK, and the only supported way to act on a ``WWW-Authenticate`` +challenge. Applications remain responsible for browser interaction, redirect +handling, consent UX, secure persistent storage, and naming the origins their +deployment may contact. -The OAuth features are primarily designed for remote transports like HTTP or WebSockets, where persistent and secure identification is required. +Lower-level pieces are also exported — PKCE generation, metadata discovery, +authorization-code exchange, refresh, volatile token storage, and +transport-level bearer authentication — but they are **components of the +manager, not an alternative to it**. See +:ref:`oauth-low-level-building-blocks` for what using them directly costs you. -OAuth Overview --------------- +Challenge-driven authorization +------------------------------ -The SDK implements the modern OAuth 2.1 authorization code flow with **PKCE** (Proof Key for Code Exchange) by default. This ensures high security for native applications and public clients without requiring client secrets. +A protected MCP server answers an unauthenticated request with ``401`` and a +``WWW-Authenticate`` header naming its protected-resource metadata (RFC 9728). +Hand that header to the manager and it runs the whole exchange: -Key components: +.. note:: -* **OAuthConfig**: Defines client identifiers and server endpoints. -* **OAuthDiscoveryClient**: Automatically resolves metadata from `.well-known` endpoints. -* **TokenStore**: Persistently stores access and refresh tokens. -* **OAuthAuthenticator**: Manages the lifecycle of tokens, including automatic background refresh. -* **OAuthClientTransport**: A transport wrapper that injects bearer tokens into outgoing requests. -* **Auth Middleware**: Server-side protection for MCP tools and resources. - -Configuration (OAuthConfig) ---------------------------- - -The `OAuthConfig` structure holds the necessary parameters for interacting with an OAuth authorization server. + The ``https://`` origins below are what a production deployment should + use. This build has no TLS support of its own — see + :ref:`oauth-security-boundary` before contacting anything but a loopback + or internally-trusted ``http://`` target. .. code-block:: cpp #include - mcp::auth::OAuthConfig config; + // A default-constructed policy refuses every origin, so name the ones this + // client may fetch metadata and tokens from, before the first request. + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {"https://mcp.example.com", "https://auth.example.com"}; + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = "https://mcp.example.com/mcp"; // also the token-store key config.client_id = "my-mcp-client"; - config.token_endpoint = "https://auth.example.com/token"; - config.authorization_endpoint = "https://auth.example.com/authorize"; - config.redirect_uri = "http://localhost:8080/callback"; - config.scope = "mcp:full_access"; + config.redirect_uri = "https://app.example.com/callback"; + config.policy = std::move(policy); + + auto token_store = std::make_shared(); -Discovery (OAuthDiscoveryClient) --------------------------------- + // The consent step. The SDK never launches a browser and never binds a + // listener for the redirect: carrying the user agent to the authorization + // endpoint and collecting the response is the application's job. + auto authorize = [](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + auto redirect_url = co_await open_in_browser(request.authorization_url); + co_return mcp::auth::parse_authorization_response(redirect_url); + }; + + auto manager = std::make_shared( + executor, token_store, std::move(config), authorize); + + if (co_await manager->try_handle_challenge(www_authenticate)) { + auto authed = std::make_shared(inner, manager); + mcp::Client client(authed, executor); + co_await client.connect("my-client", "1.0.0"); + } + +``try_handle_challenge`` returns ``false`` when the header carried no ``Bearer`` +challenge to act on, and throws when discovery is refused by the policy or the +authorization response is rejected. + +.. _oauth-what-the-manager-validates: + +What the manager validates +~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Each of these is a control you would otherwise have to write, and get right, +yourself: + +* **Metadata fetch policy.** Every discovery and token URL is checked against + :cpp:struct:`mcp::auth::MetadataFetchPolicy` before host resolution, so a + refused target is never contacted. A challenge-supplied ``resource_metadata`` + URL is attacker-influenced input. +* **Issuer binding.** The issuer is recorded from the authorization-server + metadata document the SDK itself fetched and validated, bound byte-for-byte + to the URL it came from. A redirect or a substituted document cannot move the + flow to a different issuer. +* **Cryptographic state.** The ``state`` parameter is generated from a + cryptographic random source, and a response whose ``state`` does not match the + recorded value is rejected. +* **RFC 9207 issuer validation.** Runs *before* the response's ``error``, + ``error_description`` and ``error_uri`` are read, so a response with a + mismatched issuer cannot smuggle attacker-chosen text to your user. +* **S256 PKCE.** The challenge is sent to the authorization endpoint and the + verifier is retained for the token request. +* **RFC 8707 resource indicator.** Carried into the code exchange, so the token + you receive is bound to this MCP server rather than replayable at another. + +Auditing an attempt +~~~~~~~~~~~~~~~~~~~ + +:cpp:func:`mcp::auth::OAuthAuthorizationManager::last_authorization_request` +returns the record the attempt was validated against — the state, the recorded +issuer, the resource indicator and the PKCE verifier — so an application can +audit the binding rather than take it on trust. +:cpp:func:`mcp::auth::OAuthAuthorizationManager::last_client_identity` reports +which path produced the client identity: an injected credential, a client ID +metadata document, or a dynamic registration. + +A complete runnable flow, including a mock authorization server, is in +``examples/features/oauth_flow.cpp``. + +.. _oauth-protecting-a-server: + +Protecting a server +------------------- -Instead of hardcoding endpoints, the `OAuthDiscoveryClient` can resolve authorization server details and protected resource metadata using standard `.well-known` discovery documents. +The other half of the flow above is a server that tells an unauthorized client +where to get a token. Both HTTP server transports — +:cpp:class:`mcp::StreamableHttpSessionManager` and +:cpp:class:`mcp::HttpServerTransport` — take the same three settings, and one +call is enough to make the flow work end to end: .. code-block:: cpp - auto oauth_http = std::make_shared(executor); - mcp::auth::OAuthDiscoveryClient discovery(oauth_http); + mcp::StreamableHttpSessionManager manager(executor, host, port, factory); + manager.set_bearer_token_validator( + [](std::string_view token) { return validate_token(token); }); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "https://mcp.example.com/mcp"; + metadata.authorization_servers = {"https://auth.example.com"}; + metadata.scopes_supported = {"mcp:read", "mcp:write"}; + manager.set_protected_resource_metadata(metadata); + +:cpp:func:`set_protected_resource_metadata` does two things. It serves the +RFC 9728 document, answering GET requests at its path without an +``Authorization`` header so a client holding no token can read it. And it fills +the ``resource_metadata`` parameter of the ``WWW-Authenticate`` challenge with +that document's URL, which is the value ``try_handle_challenge`` needs to begin +discovery. + +Both the path and the URL are derived from ``resource``, and from nothing else. +RFC 9728 §3.1 inserts the well-known segment between the authority and the +resource's own path, so the example above publishes at +``https://mcp.example.com/.well-known/oauth-protected-resource/mcp``. Only a +resource sitting at the origin root is described at the bare +``/.well-known/oauth-protected-resource``. Set ``metadata.path`` to override the +derivation; leave it empty to get it. + +The same derivation is available to a caller serving the document itself: +:cpp:func:`mcp::protected_resource_metadata_path` returns the path, +:cpp:func:`mcp::protected_resource_metadata_url` the absolute URL to advertise +in the challenge, and :cpp:func:`mcp::format_protected_resource_metadata` the +JSON body. Both path functions throw ``std::invalid_argument`` when ``resource`` +is not an absolute URL with an authority. + +.. important:: + + The advertised URL is never inferred from the address the transport is bound + to. Behind a TLS terminator, a reverse proxy, or a container port mapping — + which is to say in most deployments — the listener's own origin is not the one + clients can reach, so a URL built from it would send clients somewhere they + cannot fetch. ``resource`` is what the outside world calls this server, and it + is required for exactly that reason. + +To send other challenge parameters, or to point at metadata this server does not +host itself, set them explicitly: - // Discover which auth server protects a specific resource - auto resource_meta = co_await discovery.discover_protected_resource("https://api.example.com/mcp"); +.. code-block:: cpp - // Discover endpoints for an authorization server - auto auth_meta = co_await discovery.discover_auth_server(resource_meta.authorization_servers.front()); + mcp::BearerChallengeConfig challenge; + challenge.realm = "mcp"; + challenge.scope = "mcp:read"; + challenge.error = "invalid_token"; + manager.set_bearer_challenge(challenge); + +Parameters are sent in the order realm, error, scope, resource_metadata, each as +an RFC 7235 quoted-string. A ``resource_metadata`` set here is kept as written +and is not overwritten by the metadata document's URL. The two setters are +order-independent: each renders from the whole current configuration, so calling +them in either order produces the same header. A value that cannot appear in a +quoted-string is rejected by the setter with ``std::invalid_argument``, so it +cannot reach the wire as a header injection. +:cpp:func:`mcp::format_www_authenticate` renders the same string for a caller +that wants to serve the challenge from its own code. + +Validating tokens +~~~~~~~~~~~~~~~~~ + +``set_bearer_token_validator`` takes a synchronous predicate, which suits a +locally-verifiable token such as a signed JWT whose key is already in memory. +When the decision needs I/O — token introspection, a JWKS fetch — use +``set_async_bearer_token_validator`` instead, so the wait suspends rather than +blocking the executor that is concurrently serving MCP traffic: -Token Storage (TokenStore) --------------------------- +.. code-block:: cpp -Tokens are managed through the `TokenStore` interface. The SDK provides a thread-safe `InMemoryTokenStore` for development and volatile sessions. For production, you can implement a custom `TokenStore` that persists tokens to a secure database or keychain. + manager.set_async_bearer_token_validator( + [](std::string token) -> mcp::Task { + co_return co_await introspect(std::move(token)); + }); -.. code-block:: cpp +The token is passed by value because the validator may suspend, and a view into +the request buffer is not guaranteed to survive that. A transport accepts one +validator: installing both throws ``std::logic_error``. A server that installs +only the synchronous one pays nothing for the asynchronous path. - auto token_store = std::make_shared(); +The two signatures are exported as :cpp:type:`mcp::BearerTokenValidator` +(``bool(std::string_view)``) and :cpp:type:`mcp::AsyncBearerTokenValidator` +(``Task(std::string)``), so a validator can be stored and passed around +rather than written inline at the call. Neither is ever handed the +``Authorization`` header itself: the transports strip the scheme first with +:cpp:func:`mcp::http_bearer_token`, which returns the characters after +``Bearer`` and an empty view when the scheme is absent, and an empty token is +rejected before any validator runs. -PKCE Flow ---------- +Health checks and other routes +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ -For secure authorization code flows, use `generate_pkce_pair()` to create a code verifier and challenge. +Routes that must answer before a token exists are named individually: .. code-block:: cpp - auto pkce = mcp::auth::generate_pkce_pair(); - // Send pkce.code_challenge to the authorization endpoint - // Send pkce.code_verifier to the token endpoint during exchange + manager.set_unauthenticated_paths({"/health"}); -Client-Side Authentication --------------------------- +Each entry is matched for equality against the path component of the request +target, with any query string or fragment removed first, so ``/health`` also +exempts ``/health?probe=1``. :cpp:func:`mcp::http_request_path` performs that +split and is exported, so a custom request handler can key off exactly the same +path the exemption check used. -The `OAuthAuthenticator` and `OAuthClientTransport` work together to automate authentication. The transport wrapper automatically injects the `auth_token` into the `_meta` field of every MCP request. +.. warning:: -.. literalinclude:: ../../examples/features/oauth_flow.cpp - :language: cpp - :lines: 410-413 - :dedent: 8 + An exempt path is excused from the bearer check **and** excluded from MCP + dispatch. MCP is otherwise served on every path the custom handler declines, + so exempting a path without also excluding it would serve MCP there with no + authentication at all — and naming the path MCP runs on would silently + disable authentication for the whole server. An exempt request that the + metadata route and the custom request handler both decline is therefore + answered ``404 Not Found``, which makes that misconfiguration fail loudly + instead of quietly. -Automatic Token Refresh -^^^^^^^^^^^^^^^^^^^^^^^ +The request pipeline runs the protected-resource metadata route first, then the +unauthenticated path list, then the bearer check, then the custom request +handler on ``StreamableHttpSessionManager``, then MCP dispatch. Every one of +these settings is opt-in: a server that sets none of them behaves exactly as +before, down to the bare ``Bearer`` challenge on a 401. All of them must be +configured before ``listen()`` starts, and a setter called afterwards throws +``std::logic_error``. -If a request fails with an authentication error (JSON-RPC codes -32001 or -32000), the `OAuthClientTransport` will: +``examples/features/oauth_flow.cpp`` builds its challenge and its metadata route +from one ``ProtectedResourceMetadataConfig`` through these same functions, so it +shows the derivation running rather than restating its result. -1. Catch the error response. -2. Call `authenticator->try_refresh_token()` using the stored refresh token. -3. If successful, automatically replay the original request with the new token. +Request body size +~~~~~~~~~~~~~~~~~ -This ensures a seamless experience for the user even when access tokens expire. +Both transports cap the HTTP request body they will read, answering +``413 Payload Too Large`` and closing the connection beyond it. The default is +:cpp:var:`mcp::constants::g_default_max_request_body_bytes`, 8 MiB, and +``set_max_request_body_bytes`` overrides it. -Server-Side Protection ----------------------- +The default is deliberately well above Boost.Beast's own 1 MB. Binary content +reaches an MCP server as base64 inside the JSON body — ``ImageContent::data``, +``BlobResourceContents::blob`` and ``AudioContent`` have no out-of-band or +streaming path — and base64 inflates by 4/3, so a 1 MB cap admits only about +750 KB of raw bytes, which ordinary phone photos and screenshots already exceed. -On the server side, you can protect your tools and resources using `make_auth_middleware`. This middleware extracts the bearer token from incoming request metadata and validates it via a callback. +.. warning:: -.. literalinclude:: ../../examples/features/oauth_flow.cpp - :language: cpp - :lines: 524-530 - :dedent: 8 + Lower this only deliberately. A rejection happens in the HTTP parser, before + the request becomes an MCP message, so it is reported as a ``413`` on the + transport and not through the MCP error channel. A cap set near the size of + the raw content presents to users as large tool calls failing for no visible + reason. -Example Walkthrough -------------------- +.. _oauth-security-boundary: + +Security boundary +----------------- + +``OAuthHttpClient`` and the built-in MCP HTTP transports currently accept only +plain ``http://`` URLs. Use them directly only for loopback development or a +trusted internal boundary. Production traffic outside that boundary must use a +TLS-terminating proxy or a custom TLS transport. Never log access tokens, +refresh tokens, authorization codes, client secrets, or PKCE verifiers. + +.. _oauth-low-level-building-blocks: + +Low-level building blocks +------------------------- + +.. warning:: + + The pieces below perform **no** issuer binding, **no** ``state`` check and + **no** RFC 9207 validation. Composing them into your own authorization flow + compiles and appears to work, and silently gives up every control listed + under :ref:`what the manager validates `. + Use them to inspect metadata or to drive a token exchange whose issuer you + have already established by other means. To act on a ``WWW-Authenticate`` + challenge, use :cpp:class:`mcp::auth::OAuthAuthorizationManager`. -The following example demonstrates a complete OAuth integration, from discovery to authenticated tool calls with automatic refresh. +Configuration +~~~~~~~~~~~~~ -.. literalinclude:: ../../examples/features/oauth_flow.cpp - :language: cpp - :lines: 378-474 - :dedent: 0 +``OAuthConfig`` describes the endpoints and client information used for an +explicit token exchange or refresh. It is unrelated to +``OAuthAuthorizationConfig``, which configures the manager: + +.. code-block:: cpp + + #include + + mcp::auth::OAuthConfig config; + config.client_id = "my-mcp-client"; + config.token_endpoint = "http://127.0.0.1:9000/token"; + config.authorization_endpoint = "http://127.0.0.1:9000/authorize"; + config.redirect_uri = "http://127.0.0.1:8080/callback"; + config.scope = "mcp:read"; + +Discovery and PKCE +~~~~~~~~~~~~~~~~~~ + +``OAuthDiscoveryClient`` reads protected-resource and authorization-server +metadata. Discovery results are cached for a configurable interval. Install a +:cpp:struct:`mcp::auth::MetadataFetchPolicy` on the ``OAuthHttpClient`` first, +or every fetch is refused. + +.. code-block:: cpp + + auto oauth_http = std::make_shared(executor); + mcp::auth::OAuthDiscoveryClient discovery(oauth_http); + + auto resource = co_await discovery.discover_protected_resource(server_url); + auto authorization_server = + co_await discovery.discover_auth_server(resource.authorization_servers.front()); + auto pkce = mcp::auth::generate_pkce_pair(); + +Taken this far by hand, the application must then build and open the +authorization URL, receive the callback, generate and check its own ``state``, +apply RFC 9207 issuer validation itself, and call ``exchange_code`` with the +original PKCE verifier. The manager does all of that, correctly, and is the +reason this route is not recommended. + +Token storage and refresh +~~~~~~~~~~~~~~~~~~~~~~~~~ + +``InMemoryTokenStore`` is thread-safe but volatile. Production applications +should implement ``TokenStore`` with encryption and operating-system-appropriate +secret storage. + +``OAuthAuthenticator`` reads the current access token and can explicitly +refresh it, for an application that already holds a token it obtained some +other way. ``OAuthAuthorizationManager`` implements the same +:cpp:class:`mcp::auth::Authenticator` interface, so either can be handed to +``OAuthClientTransport``; only the manager can acquire a token in the first +place. Neither runs a background refresh task. ``OAuthClientTransport`` +attempts one refresh-and-retry after an HTTP 401 from ``HttpClientTransport``; +for non-HTTP transports it also understands the legacy JSON-RPC +``g_UNAUTHORIZED`` path. Because a request can be replayed, callers should +avoid wrapping non-idempotent operations unless their server provides its own +deduplication semantics. + +HTTP authentication +~~~~~~~~~~~~~~~~~~~ + +When ``OAuthClientTransport`` wraps ``HttpClientTransport``, the access token +is sent as ``Authorization: Bearer ``. It is not copied into MCP request +metadata. Configure the matching server-side validator before starting the +listener: + +.. code-block:: cpp + + mcp::StreamableHttpSessionManager manager(executor, host, port, factory); + manager.set_bearer_token_validator( + [](std::string_view token) { return validate_token(token); }); + +``HttpServerTransport`` exposes the same validator. The validator decides +whether a token is acceptable; what an unauthorized client is told to do next is +the challenge, covered under :ref:`oauth-protecting-a-server`. The older +``make_auth_middleware`` helper checks ``_meta.auth_token`` and exists for +non-HTTP or legacy integrations; it should not replace authentication at the +HTTP boundary. + +Example and current limits +-------------------------- -Cross-References +``examples/features/oauth_flow.cpp`` is a loopback demonstration using a mock +authorization server and ``MemoryTransport`` for MCP messages. It drives +``OAuthAuthorizationManager`` from a ``WWW-Authenticate`` challenge through +discovery, consent, RFC 9207 validation and a resource-indicated code exchange, +then performs one refresh, without exposing secret values. Its consent callback +mints the code directly from the mock authorization server, which is possible +only because the example owns both ends; the file says so at that point and +should not be copied there. + +Its MCP-side authentication uses the legacy ``_meta.auth_token`` middleware +path, so it does not demonstrate the HTTP ``Authorization: Bearer`` boundary +described above. + +Both ends of the flow are in the SDK. A server configured as +:ref:`oauth-protecting-a-server` describes answers an unauthenticated request +with a challenge naming its protected-resource metadata, serves that metadata +without a token, and ``OAuthAuthorizationManager`` drives the rest from that +header. The server side is deliberately narrow: it advertises a challenge and +publishes a metadata document, and it neither issues tokens nor validates them +beyond calling the ``BearerTokenValidator`` the application supplies, so an +authorization server and the token verification behind that callback remain the +deployment's own. + +The official client conformance baseline records the OAuth scenarios that are +not yet implemented. In particular, the SDK does not currently orchestrate all +protected-resource metadata variants, CIMD/pre-registration, scope escalation, +or token-endpoint authentication modes as a complete end-user flow. + +Cross-references ---------------- -* See :doc:`transports` for more information on transport wrappers. -* See :doc:`context` for details on request metadata. +* See :doc:`transports` for HTTP transport security and Origin validation. +* See :doc:`context` for request context and the legacy metadata path. diff --git a/docs/concepts/pagination.rst b/docs/concepts/pagination.rst index 7dff408..7ab7ecd 100644 --- a/docs/concepts/pagination.rst +++ b/docs/concepts/pagination.rst @@ -10,9 +10,11 @@ Server-side Configuration Servers can control the maximum number of items returned in a single list response by calling ``set_page_size()``. -.. literalinclude:: ../../include/mcp/server/server.hpp - :language: cpp - :lines: 263-268 +.. code-block:: cpp + + // Return at most 100 entries from each tools/list, resources/list, + // resources/templates/list, or prompts/list request. + server.set_page_size(100); When ``set_page_size`` is set to a non-zero value, any list request (e.g., ``tools/list``, ``resources/list``) that exceeds this limit will automatically be paginated. The server will return a ``nextCursor`` in the result, which the client can use to fetch the next page. @@ -62,11 +64,8 @@ The SDK handles the complexity of slicing the data and generating cursors automa Implementation Details ---------------------- -Internal pagination logic uses the ``paginate`` helper and ``PaginationSlice`` structure: - -.. literalinclude:: ../../include/mcp/server/server.hpp - :language: cpp - :lines: 449-455 +Internal pagination logic validates the cursor, calculates the current slice, +and returns an opaque next cursor only when another page remains. The process of pagination is fully transparent to the server developer once ``set_page_size()`` is called. The SDK internal dispatch logic detects when a list request arrives, calculates the appropriate slice of items based on the provided cursor (or lack thereof), and constructs the response with the ``nextCursor`` if more items remain. This ensures that even as the number of tools or resources grows, the memory footprint and network payload per request remain stable. diff --git a/docs/concepts/prompts.rst b/docs/concepts/prompts.rst index 7a99ff5..53396b6 100644 --- a/docs/concepts/prompts.rst +++ b/docs/concepts/prompts.rst @@ -14,18 +14,32 @@ specific arguments to get a list of messages to include in the conversation. Registering a Basic Prompt -------------------------- -A basic prompt can be registered using the `add_prompt` method. The simplest -overload takes a name, description, and a set of arguments (as JSON schema) -along with a synchronous handler. This method is ideal for simple use cases -where you don't need full type safety for the handler parameters or result. +Prompts are registered with `add_prompt`, which is a template over the handler's +input and output types: ``add_prompt(prompt, handler)``. Passing +`nlohmann::json` for both is the untyped route, and suits simple cases where you +don't need full type safety for the handler parameters or result. .. literalinclude:: ../../examples/servers/stdio/server_stdio.cpp :language: cpp - :lines: 125-132 + :start-after: docs-begin: untyped-prompt + :end-before: docs-end: untyped-prompt :dedent: 8 -The handler receives a `nlohmann::json` object containing the arguments and -returns a `nlohmann::json` object representing the `GetPromptResult`. +.. warning:: + + A prompt handler receives the **entire** ``GetPromptRequestParams``, not just + the arguments. For a ``prompts/get`` request naming ``greet`` with argument + ``who``, the handler is passed ``{"name": "greet", "arguments": {"who": + "world"}}``, so the argument is read as + ``params.value("arguments", json::object()).value("who", ...)``. Reading + ``params.value("who", ...)`` directly compiles, raises no error, and silently + yields the default value. + + This differs from tool handlers, which receive the tool arguments already + unwrapped. A prompt carries a name its handler may need, so the whole + parameter object is passed through. + +The handler returns a `nlohmann::json` object representing the `GetPromptResult`. Parameterized Prompts and Arguments ----------------------------------- diff --git a/docs/concepts/resources.rst b/docs/concepts/resources.rst index ab3b853..19bd910 100644 --- a/docs/concepts/resources.rst +++ b/docs/concepts/resources.rst @@ -16,17 +16,28 @@ A static resource is a direct mapping from a fixed URI to a specific value or handler. These are ideal for configuration files, documentation, or fixed data sets that don't change their identification scheme. -You can add a static resource using the ``add_resource`` method. The library -provides high-level overloads for common use cases: +Register metadata and a typed read handler with ``add_resource``: .. code-block:: cpp #include - server.add_resource("mcp://logs/system", "System logs", "text/plain", - []() -> std::string { - return "System is running normally."; - }); + mcp::Resource resource; + resource.uri = "mcp://logs/system"; + resource.name = "System logs"; + resource.mimeType = "text/plain"; + + server.add_resource( + resource, [](mcp::ReadResourceRequestParams request) { + mcp::TextResourceContents contents; + contents.uri = std::move(request.uri); + contents.mimeType = "text/plain"; + contents.text = "System is running normally."; + + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(contents)); + return result; + }); For more complex resources, you can use the structured protocol types: @@ -39,21 +50,30 @@ For more complex resources, you can use the structured protocol types: Resource Templates ------------------ -Resource templates allow you to define a URI pattern with parameters using the -RFC 6570 URI Template syntax. This is powerful for exposing collections of data -where the specific resource identity depends on a parameter, such as a database -record ID or a filename. - -The server will expand the URI when a client requests it, and pass the extracted -parameters to your handler. +Resource templates let one handler serve a family of URIs. The current matcher +supports the common RFC 6570 expression forms for routing, but it is not a full +RFC 6570 expansion or variable-extraction engine. The handler receives the full +requested URI and can parse or validate application-specific segments itself. .. code-block:: cpp - server.add_resource_template("mcp://logs/{node}", "Node logs", - [](const std::map& params) -> std::string { - std::string node = params.at("node"); - return "Logs for node: " + node; - }); + mcp::ResourceTemplate logs; + logs.uriTemplate = "mcp://logs/{node}"; + logs.name = "Node logs"; + logs.mimeType = "text/plain"; + + server.add_resource_template( + logs, [](mcp::ReadResourceRequestParams request) { + mcp::TextResourceContents contents; + contents.uri = std::move(request.uri); + contents.mimeType = "text/plain"; + contents.text = "Logs for requested node"; + + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(contents)); + return result; + }); Templates are particularly useful when you have a large or open-ended set of resources that share the same schema or purpose. @@ -95,17 +115,19 @@ For more details on how notifications work across the protocol, see :doc:`/index Resource Change Notifications ----------------------------- -When a resource's content changes, the server should notify all active -subscribers using ``notify_resource_updated``. The library handles the -routing of these notifications to the correct clients. +When a resource's content changes, call ``notify_resource_updated`` on the +``Server`` instance for the relevant session. The method sends only when that +session subscribed to the exact URI. In a multi-session deployment, the +application is responsible for invoking the notification on each relevant +per-session server instance. .. code-block:: cpp // Notify subscribers that a specific resource has changed co_await server.notify_resource_updated("mcp://status/counter"); -This mechanism ensures that the LLM or the client application always has -access to the most up-to-date information without constant polling. +The notification tells a subscribed client that it should read the resource +again; delivery and refresh timing still depend on the connection and client. .. literalinclude:: ../../examples/features/notifications_subscriptions.cpp :language: cpp @@ -137,6 +159,7 @@ Best Practices interpret the data. - **Error Handling**: Throwing an exception in a resource handler will automatically return an appropriate JSON-RPC error to the client. -- **Context Awareness**: Use the handler's ``Context`` object for logging or - reporting progress for long-running resource generation. See :doc:`context` - for more information. +- **Context Awareness**: A resource handler can use ``Context`` for logging and + reverse requests. Resource reads do not currently propagate a request + ``progressToken`` into the handler context, so do not rely on progress + reporting for them. See :doc:`context` for more information. diff --git a/docs/concepts/roots.rst b/docs/concepts/roots.rst index b7f6d3d..7c7625a 100644 --- a/docs/concepts/roots.rst +++ b/docs/concepts/roots.rst @@ -14,14 +14,17 @@ In the Model Context Protocol, roots represent the "home" or "workspace" directo Client-side Setup ----------------- -A client provides its list of roots using the ``set_roots()`` method. This method stores the roots internally and automatically registers a handler for the ``roots/list`` request. +A client provides its list of roots using the ``set_roots()`` method. This method stores the roots internally, registers a handler for ``roots/list``, and—when called before ``connect()``—automatically advertises the matching client capability. .. literalinclude:: ../../examples/features/roots.cpp :language: cpp :start-after: // Declare roots that the client can access :end-before: Implementation client_info; -The ``set_roots`` method also takes an optional ``notify`` parameter. If set to ``true``, the client will send a ``notifications/roots/list_changed`` notification to the server whenever the roots are updated, provided the server supports this capability. +The ``set_roots`` method also takes an optional ``notify`` parameter. If set to +``true``, the client sends ``notifications/roots/list_changed`` only when its +initialization capabilities advertised ``roots.listChanged=true``. This is a +client capability; servers do not advertise roots support. Server-side Discovery --------------------- diff --git a/docs/concepts/tools.rst b/docs/concepts/tools.rst index a9284fc..a2ac49d 100644 --- a/docs/concepts/tools.rst +++ b/docs/concepts/tools.rst @@ -113,14 +113,12 @@ Error Handling There are two primary ways to handle errors in tool handlers: -1. **Throwing Exceptions**: The synchronous ``add_tool(name, description, - schema, std::function)`` overload - catches handler exceptions and converts them into a tool error response - (``isError: true``). Typed and asynchronous handlers propagate exceptions as - JSON-RPC internal errors instead. -2. **Explicit Error Returns**: You can return a `mcp::CallToolResult` object - with the `isError` field set to `true`. This gives you full control over - the error message and any additional metadata you want to return. +1. **Throwing Exceptions**: All tool-handler overloads convert standard + exceptions into a tool error response (``isError: true``). +2. **Explicit Error Returns**: A handler declared with + ``Out=mcp::CallToolResult`` can return a complete result whose ``isError`` + field is ``true``. ``add_raw_tool()`` provides the equivalent low-level JSON + escape hatch with runtime validation. The following example demonstrates both patterns: @@ -129,10 +127,10 @@ The following example demonstrates both patterns: :lines: 45-76 Use ``isError: true`` when the tool completed normally but needs to report an -application-level failure back to the client. Unhandled exceptions from typed or -async handlers become JSON-RPC errors, while exceptions from the raw synchronous -JSON overload are converted into tool error results. In both cases, the SDK -keeps the session alive when possible. This is a critical distinction: +application-level failure back to the client. Standard exceptions from sync, +typed, and async tool handlers are converted into the same protocol-level tool +error shape. In both cases, the SDK keeps the session alive when possible. This +is a critical distinction: - **Protocol Errors**: Issues like malformed JSON, invalid method names, or disconnected transports cause JSON-RPC level errors. diff --git a/docs/concepts/transports.rst b/docs/concepts/transports.rst index ed3f1a1..1404e55 100644 --- a/docs/concepts/transports.rst +++ b/docs/concepts/transports.rst @@ -16,11 +16,12 @@ The SDK includes several built-in transport types to cover different use cases: 1. **StdioTransport**: For local process communication via standard input/output. This is the most common choice for connecting to local agents like Claude Desktop. -2. **HttpServerTransport / HttpClientTransport**: For MCP over HTTP with custom - streamable headers, following the official MCP specification. -3. **WebSocketServerTransport / WebSocketClientTransport**: For persistent, - bidirectional network communication. - Ideal for high-performance remote integrations. +2. **HttpServerTransport / HttpClientTransport**: Single-session Streamable HTTP + building blocks. ``StreamableHttpSessionManager`` provides a multi-session + server endpoint. +3. **WebSocketServerTransport / WebSocketClientTransport**: An optional + persistent transport for integrations that explicitly agree on WebSocket; + it is not the standard Streamable HTTP transport. 4. **MemoryTransport**: An in-memory transport for unit testing and in-process communication. @@ -78,8 +79,21 @@ Clients can connect to a stdio-based server by providing the executor to the HTTP Transport (Network) ------------------------ -The HTTP transport follows the MCP over HTTP specification, which uses a long-running -POST request for server-sent events or streamable data. +The HTTP implementation accepts MCP POST requests and returns either JSON or a +finite SSE body. Managed sessions support session IDs, DELETE teardown, and +bounded event replay through GET with ``Last-Event-ID``. Continuously open SSE +polling streams are not yet implemented. + +After initialization, ``HttpClientTransport`` carries the server-selected +protocol version on subsequent session requests and best-effort DELETE +teardown, including when the initialize response arrived in an SSE event. + +The built-in HTTP transports are plaintext. Bind local development servers to +loopback. For remote deployments, place them behind a TLS-terminating proxy or +provide a custom TLS transport. Requests carrying an ``Origin`` header are +denied by default; configure an exact allowlist (or explicitly opt into all +origins) before starting the listener. Bearer-token validators must likewise be +installed before ``listen()`` or ``run()``. HTTP Server Convenience ~~~~~~~~~~~~~~~~~~~~~~~ @@ -87,6 +101,13 @@ HTTP Server Convenience For developers who want a quick way to host an MCP server over HTTP without worrying about the underlying networking boilerplate, the SDK provides `run_http()`. +``Server::run_http()`` owns its ``HttpServerTransport`` internally and does not +currently expose Origin allowlist or bearer-token validator configuration. Use +it only on a suitable trusted boundary, such as the loopback example below. For +a configurable deployment, construct ``HttpServerTransport`` or +``StreamableHttpSessionManager`` directly, apply the security settings, and +then start the listener. + .. literalinclude:: ../../examples/features/http_server_convenience.cpp :language: cpp :start-after: // This convenience method handles everything: @@ -115,10 +136,11 @@ yourself. WebSocket Transport ------------------- -For full-duplex, bidirectional communication over the network, +For integrations that explicitly choose full-duplex WebSocket communication, ``WebSocketServerTransport`` and ``WebSocketClientTransport`` -are the ideal choices. Unlike HTTP, which often requires polling or long-running -streams for bidirectional data, WebSockets provide a native persistent connection. +provide a persistent connection. Interoperability with standard MCP clients is +not implied; use Streamable HTTP when protocol-standard remote transport is +required. .. code-block:: cpp @@ -187,17 +209,17 @@ integration. Use the table below as a guide: | | | - Secure by default (local only) | | | | - Easiest to deploy | +----------------------+--------------------------+------------------------------------------------+ -| **HTTP** | Remote APIs, Web Hooks | - Firewall and proxy friendly | -| | | - Works with standard web infrastructure | -| | | - Good for one-way or slow bidirectional data | +| **HTTP** | MCP remote endpoints | - Standard Streamable HTTP request model | +| | | - Managed sessions and bounded event replay | +| | | - Deploy behind TLS for non-loopback use | +----------------------+--------------------------+------------------------------------------------+ -| **WebSocket** | High-performance Remote | - Native bidirectional communication | -| | | - Persistent, low-latency connection | -| | | - Minimal message overhead | +| **WebSocket** | Agreed custom integration| - Full-duplex persistent channel | +| | | - Separate client and server transport types | +| | | - Requires explicit interoperability agreement | +----------------------+--------------------------+------------------------------------------------+ -| **Memory** | Unit & Integration Tests | - Extremely fast execution | -| | | - Deterministic (no network flakes) | -| | | - No side effects on the host system | +| **Memory** | Unit & Integration Tests | - In-process message exchange | +| | | - No network setup | +| | | - Linked endpoint-pair helper | +----------------------+--------------------------+------------------------------------------------+ Custom Transports diff --git a/docs/conf.py b/docs/conf.py index be387fc..9e1c84a 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -5,6 +5,8 @@ import sys from pathlib import Path +from docutils.parsers.rst import Directive, directives + # -- Project information ----------------------------------------------------- project = 'mcp-cpp-sdk' copyright = '2026, MCP C++ SDK Contributors' @@ -63,8 +65,6 @@ # directories to ignore when looking for source files. # Exclude API reference pages when Doxygen XML is not available. exclude_patterns = ['_build', 'Thumbs.db', '.DS_Store'] -if not _has_doxygen: - exclude_patterns.append('api') # The name of the Pygments (syntax highlighting) style to use. pygments_style = 'sphinx' @@ -102,3 +102,23 @@ # Output file base name for HTML help builder. htmlhelp_basename = 'mcp-cpp-sdkdoc' + + +class _UnavailableDoxygenDirective(Directive): + """Keep the API overview buildable when Doxygen XML is unavailable.""" + + required_arguments = 1 + option_spec = { + 'members': directives.flag, + 'undoc-members': directives.flag, + } + + def run(self): + return [] + + +def setup(app): + """Install inert Doxygen directives for the documented fallback build.""" + if not _has_doxygen: + for name in ('doxygenclass', 'doxygenstruct', 'doxygentypedef'): + app.add_directive(name, _UnavailableDoxygenDirective) diff --git a/docs/contributing.rst b/docs/contributing.rst index 802a3a8..ce879dc 100644 --- a/docs/contributing.rst +++ b/docs/contributing.rst @@ -149,6 +149,9 @@ Running Tests # Run the GoogleTest binary directly with a filter ./build/mcp-sdk-tests --gtest_filter=ServerCoreTest.* +A test build from a checkout needs Python 3.9+ for the peer-input matrix check; +opt out with ``-DMCP_CPP_SDK_CHECK_JSON_MATRIX=OFF``. + Code Coverage ^^^^^^^^^^^^^ @@ -162,6 +165,41 @@ Aim for >80% code coverage for new features: # View report open build/coverage/index.html +Sanitizers +^^^^^^^^^^ + +CI runs the test suite under each sanitizer on every push. The same builds run +locally: + +.. code-block:: bash + + # AddressSanitizer + UndefinedBehaviorSanitizer + python scripts/build.py --sanitize --test + + # ThreadSanitizer + python scripts/build.py --tsan --test + + # Either of the two, compiled with Clang + python scripts/build.py --tsan --compiler clang --test + + # MemorySanitizer (Clang only) + python scripts/build.py --msan --test + +Run ThreadSanitizer for code that runs on several threads or is closed from +another thread. The two reports that are not defects (Asio's fence-based +reference count and its signal handler) are suppressed in ``test/tsan.supp``, +each with its reason. + +MemorySanitizer reports a read of any memory that uninstrumented code wrote, +so ``--msan`` builds its own libc++ from a pinned LLVM commit into +``build/msan-libcxx`` on first use, and has Conan rebuild Boost, OpenSSL and +GoogleTest against it. The first run is slow; later runs reuse both. It needs +``clang-18`` and ``git``. + +The ThreadSanitizer and MemorySanitizer builds run under ``setarch -R`` where +it is available, because neither can start on kernels with a high ASLR +entropy. + Pull Request Process -------------------- @@ -300,7 +338,7 @@ Example bug report: **Steps to Reproduce** ```cpp mcp::Client client(transport, executor); - co_await client.call_tool("nonexistent", {}); // Hangs forever + co_await client.call_tool("nonexistent", nlohmann::json::object()); // Hangs forever ``` **Expected**: Exception thrown or timeout @@ -384,6 +422,50 @@ callee's suspensions. The crash manifests as ``munmap_chunk(): invalid pointer`` at process exit, not at the point of access, making it hard to debug without valgrind. +GCC 12 and 13 Initializer-List Coroutine ICE +--------------------------------------------- + +**Never write an initializer list inside a co_await expression when its elements +have non-trivial destructors.** GCC 12 and GCC 13 crash outright: + +.. code-block:: text + + internal compiler error: in build_special_member_call, at cp/call.cc:11096 + +Unlike the GCC 11 bug above, this is a compile-time failure, not a runtime one, and the +diagnostic points at the closing brace of the enclosing lambda rather than the offending +argument. GCC 13 is the default compiler on Ubuntu 24.04. + +**Affected** (each crashes GCC 12 and 13): + +.. code-block:: cpp + + co_await client.call_tool("hello", nlohmann::json{{"name", "World"}}); + co_await client.call_tool("hello", nlohmann::json({{"name", "World"}})); // parens do not help + co_await client.get_prompt("greet", std::map{{"who", "you"}}); + co_await client.call_tool("e", std::vector{"a", "b"}); + +**Not affected**: ``nlohmann::json::object()``, ``std::vector{1, 2, 3}`` (trivially +destructible elements), a ``std::string`` temporary, and aggregate initialization such as +``mcp::ClientCapabilities{}`` or ``AddArgs{.a = 3.0, .b = 4.0}``. The trigger is the +initializer-list backing array, not the type. + +**Fix pattern — hoist-before-await**: bind the value to a named variable before the +co_await expression. + +.. code-block:: cpp + + // DO NOT: co_await client.call_tool("hello", nlohmann::json{{"name", "World"}}); + nlohmann::json arguments{{"name", "World"}}; + auto result = co_await client.call_tool("hello", arguments); + +Changing the SDK's own signatures does not help: a parameter taken by value instead of by +const reference still crashes. Avoiding the construct at the call site is the only fix. + +**Detection**: ``scripts/check_readme_snippets.py`` compiles the README's complete +examples with the toolchain that built the SDK, and CI runs it on the Linux matrix +entries, where that toolchain is GCC. + Feature Requests ---------------- diff --git a/docs/examples.rst b/docs/examples.rst index 3371263..ea1661a 100644 --- a/docs/examples.rst +++ b/docs/examples.rst @@ -8,22 +8,24 @@ directory. Overview -------- -The SDK provides seven core examples plus 12 specialized feature examples: +The SDK builds nine core examples plus 12 specialized feature examples: 1. **server_stdio** - Full-featured MCP server over stdio -2. **client_stdio** - MCP client demonstrating all client operations -3. **server_with_sampling** - Server with reverse RPC (sampling) -4. **echo_websocket** - WebSocket transport example -5. **http_loopback** - HTTP transport loopback demo -6. **llama_server** - llama.cpp MCP adapter -7. **debugger_server** - LLDB debugger MCP server +2. **server_simple** - Minimal MCP server over stdio +3. **client_stdio** - MCP client demonstrating representative server operations +4. **interactive_client** - Interactive stdio client +5. **server_with_sampling** - Server with reverse RPC (sampling) +6. **echo_websocket** - WebSocket transport example +7. **http_loopback** - HTTP transport loopback demo +8. **llama_server** - llama.cpp MCP adapter +9. **debugger_server** - LLDB debugger MCP server (uses stubs when LLDB is unavailable) Feature Examples ---------------- -The ``examples/features/`` directory contains 12 specialized examples, each -demonstrating a specific MCP capability using the in-process loopback pattern -for fast, deterministic testing. +The ``examples/features/`` directory contains 12 specialized examples. They use +small, self-contained setups to demonstrate individual capabilities; some use +in-memory loopback and others exercise local network transports. 1. **progress_cancellation** - Real-time progress updates and request cancellation 2. **notifications_subscriptions** - Server-to-client notifications and subscription handling @@ -36,7 +38,7 @@ for fast, deterministic testing. 9. **transport_memory** - Using in-memory transports for testing and modularity 10. **graceful_shutdown** - Clean exit patterns and signal handling 11. **http_server_convenience** - Simplified HTTP server setup via ``Server::run_http()`` -12. **oauth_flow** - Full OAuth 2.0 authentication flow (PKCE, token refresh) +12. **oauth_flow** - Loopback OAuth building blocks (discovery, PKCE, exchange, refresh) These examples are non-interactive and designed to be run as part of a test suite or to understand specific API patterns. @@ -44,10 +46,9 @@ test suite or to understand specific API patterns. Benchmarks ---------- -Performance-focused examples demonstrating the efficiency of different transport -layers under load. These are quick teaching benchmarks that run locally and in -CI; for sustained multi-service load testing, use the separate ``benchmark/`` -suite. +These examples exercise different transport paths and report measurements from +the local run. They are teaching benchmarks that run locally and in CI; for +sustained multi-service load testing, use the separate ``benchmark/`` suite. 1. **benchmark_stdio** - Measures roundtrip latency and throughput over the in-memory stdio-equivalent transport 2. **benchmark_http** - Performance analysis of the HTTP transport layer @@ -56,10 +57,11 @@ suite. Each benchmark performs a fixed number of iterations and reports total time, average latency, and calls per second. The current examples run 1000 iterations for the memory and HTTP transports, and 50 iterations for the WebSocket -transport so the example remains fast and deterministic in automation. +transport to keep the automated run bounded. -Each example is fully functional and can be used as a starting point for your -own MCP applications. +Use the example closest to your integration as an API reference, then add the +deployment-specific validation, security, and error handling your application +requires. server_stdio ------------ @@ -70,19 +72,19 @@ A complete MCP server implementation demonstrating: * **All 4 handler signatures**: Async with context, async without context, sync with context, sync without context * **Custom types**: JsonSerializable structs with to_json/from_json -* **Tools**: Mathematical operations (add, multiply) with typed parameters -* **Resources**: Static and dynamic resource handlers -* **Resource Templates**: URI templates with variable substitution +* **Tools**: Typed add/greet handlers and JSON echo/uppercase handlers +* **Resources**: A typed static resource handler +* **Resource Templates**: Metadata registration for a URI template * **Prompts**: Example prompt with arguments * **Context logging**: Using ``Context::log_info()`` and other log levels -This is the most comprehensive example, showing virtually every SDK feature. +This example combines the main server registration patterns in one program. **Key patterns demonstrated**: * Type-safe tool handlers with custom input/output types * Resource listing and reading -* Resource template URI expansion +* Resource template metadata listing * Prompt registration and retrieval * Structured logging via Context @@ -98,24 +100,22 @@ client_stdio **Location**: ``examples/clients/stdio/client_stdio.cpp`` -A complete MCP client implementation demonstrating: +A stdio MCP client demonstrating: * **Connection**: Initialize handshake with server capabilities * **Tools**: List available tools, call tools with typed arguments * **Resources**: List resources, read resource contents -* **Resource Templates**: List templates, expand URIs +* **Resource Templates**: List template metadata * **Prompts**: List prompts, get prompt with arguments * **Completion**: Request completion suggestions * **Notifications**: Send notifications to server -* **Graceful shutdown**: Proper connection close **Key patterns demonstrated**: * Client initialization and capability negotiation * Calling server methods with typed parameters * Handling server responses -* Error handling with try/catch -* Clean resource cleanup +* Handling typed content variants in server responses **Build and run**: @@ -140,10 +140,10 @@ enabling bidirectional communication patterns. **Key patterns demonstrated**: -* Declaring sampling capability in ServerCapabilities +* Relying on the connected client to advertise its sampling capability * Using ``Context::sample_llm()`` for reverse RPC * Awaiting client responses within a tool handler -* Structured error handling for sampling failures +* Reading text content from the sampling response **Build and run**: @@ -220,8 +220,8 @@ MCP adapter for llama.cpp's ``llama-server`` (chat, completion, embedding): * **Dual transport support**: Runtime selection between stdio and HTTP transports * **CLI argument parsing**: Comprehensive configuration via command line flags -This example demonstrates how to wrap an existing external service into -a fully compliant MCP server. +This example demonstrates how to wrap an existing external service in an MCP +server with stdio and HTTP deployment options. **Key patterns demonstrated**: diff --git a/docs/getting-started.rst b/docs/getting-started.rst index 6e90229..c1ac575 100644 --- a/docs/getting-started.rst +++ b/docs/getting-started.rst @@ -159,11 +159,11 @@ with other async work — you can wire up the transport manually: }); boost::asio::io_context io; - auto transport = std::make_unique(io.get_executor()); + auto transport = std::make_shared(io.get_executor()); boost::asio::co_spawn( io, [&]() -> mcp::Task { - co_await server.run(std::move(transport), io.get_executor()); + co_await server.run(transport, io.get_executor()); }, boost::asio::detached); io.run(); @@ -177,12 +177,13 @@ Quick Start: Minimal Client A convenience ``client.run_stdio()`` API is planned. For now, clients require manual ``io_context`` setup as shown below. -Here's a minimal MCP client that connects to a server and calls a tool: +Here's a minimal MCP client that connects to an already running Streamable HTTP +server and calls a tool: .. code-block:: cpp #include - #include + #include #include #include #include @@ -193,8 +194,9 @@ Here's a minimal MCP client that connects to a server and calls a tool: boost::asio::io_context io; - auto transport = std::make_unique(io.get_executor()); - Client client(std::move(transport), io.get_executor()); + auto transport = std::make_shared( + io.get_executor(), "http://127.0.0.1:3000/mcp"); + Client client(transport, io.get_executor()); boost::asio::co_spawn( io, @@ -203,14 +205,14 @@ Here's a minimal MCP client that connects to a server and calls a tool: Implementation info; info.name = "math-client"; info.version = "1.0.0"; - co_await client.connect(std::move(info), {}); + co_await client.connect(info, {}); // Call the "add" tool nlohmann::json args = {{"a", 5}, {"b", 7}}; auto result = co_await client.call_tool("add", args); - std::cout << "Result: " << result.dump(2) << std::endl; + std::cout << "Result: " << nlohmann::json(result).dump(2) << std::endl; - co_await client.close(); + client.close(); }(), boost::asio::detached ); @@ -218,6 +220,25 @@ Here's a minimal MCP client that connects to a server and calls a tool: io.run(); } +.. important:: + + **Client requests time out after 30 seconds by default.** The ``Client`` + above takes no options, so every ``call_tool``, ``read_resource`` and + ``list_*`` it issues carries a 30-second deadline and throws + ``mcp::McpError`` with code ``-32001`` when it expires — even though the + server may still be working on the call. Set + ``ClientOptions::request_timeout`` before constructing the client if any of + your tools legitimately run longer: + + .. code-block:: cpp + + mcp::ClientOptions options; + options.request_timeout = std::chrono::minutes(5); + Client client(transport, io.get_executor(), options); + + See :doc:`guides/error-handling` for per-request deadlines and what a + timeout does and does not guarantee about the server. + Next Steps ---------- diff --git a/docs/guides/error-handling.rst b/docs/guides/error-handling.rst index 3188ad2..e1046c0 100644 --- a/docs/guides/error-handling.rst +++ b/docs/guides/error-handling.rst @@ -36,7 +36,7 @@ Server-side Error Handling Server handlers (tools, resources, prompts) can report errors in two primary ways: -1. **Throwing Exceptions**: Exceptions from typed or asynchronous handlers propagate as JSON-RPC ``Internal error (-32603)`` responses. The raw synchronous ``add_tool(name, description, schema, std::function)`` overload is different: it catches exceptions and converts them into tool results with ``isError: true``. +1. **Throwing Exceptions**: Exceptions from a tool handler — typed, asynchronous or raw — are converted into tool results with ``isError: true``, because a tool that cannot complete has failed at its own task rather than at the protocol. Exceptions from resource and prompt handlers, which have no such result channel, propagate as JSON-RPC ``Internal error (-32603)`` responses, and so do exceptions from middleware: middleware decides whether a call may proceed at all, so refusing one is a protocol-level answer rather than a tool outcome. Exception messages are flattened and length-bounded before they reach the peer. 2. **Returning Error Results**: For application-level errors (e.g., "File not found" or "Invalid input"), handlers can return a ``CallToolResult`` with the ``isError`` flag set to ``true``. This allows the client to distinguish between a technical failure (like a crash or timeout) and a logical error within the tool's execution. Choosing Between Exceptions and Error Results @@ -87,6 +87,65 @@ Clients should always wrap server calls in ``try-catch`` blocks to handle potent :start-after: // ========== TEST 1: Catch exception from thrown error ========== :end-before: // Wait between tests +Request Timeouts +~~~~~~~~~~~~~~~~ + +.. important:: + + **Every client request carries a 30-second deadline by default.** Nothing + opts into it: ``ClientOptions::request_timeout`` starts at + ``std::chrono::seconds(30)``, and a default-constructed ``Client`` uses it + for ``call_tool``, ``read_resource``, ``list_tools`` and every other + request. A tool that legitimately runs longer than that — a large model + call, a slow build, a batch job — fails on the client side while the server + is still working on it. + +Raise it for the whole client by passing ``ClientOptions`` to the constructor: + +.. code-block:: cpp + + mcp::ClientOptions options; + options.request_timeout = std::chrono::minutes(5); + mcp::Client client(transport, io.get_executor(), options); + +A single request can be given its own deadline through ``RequestOptions``, but +only on the untyped ``send_request`` overload. The typed helpers — +``call_tool``, ``read_resource``, ``list_tools`` and the rest — take no +per-request options and always use the client-wide default, so a client that +makes one slow call among many fast ones either raises the default for all of +them or drops to ``send_request`` for that one: + +.. code-block:: cpp + + mcp::CallToolParams call_params; + call_params.name = "render"; + call_params.arguments = args; + + mcp::RequestOptions slow; + slow.timeout = std::chrono::minutes(30); + + auto raw = co_await client.send_request( + "tools/call", nlohmann::json(call_params), slow); + auto result = raw.get(); + +When the deadline expires the awaiting coroutine throws +:cpp:class:`mcp::McpError` with code ``-32001`` +(:cpp:var:`mcp::g_REQUEST_TIMEOUT`) and the message ``Request timed out``. +Catch it the same way as any other ``McpError``. + +.. warning:: + + A timeout is a local decision, not a cancellation. A request whose bytes + have not been written yet is dropped before transmission, but once transport + I/O has started the server may still run the call to completion after the + client has given up. Treat a timed-out non-idempotent request as having an + unknown outcome, and see the retry guidance below before repeating it. + +A non-positive timeout is rejected rather than treated as "no deadline": +constructing a ``Client`` with one, or passing one in ``RequestOptions``, +throws ``std::invalid_argument``. The SDK offers no way to disable the deadline +entirely — use a duration long enough for the slowest call you expect. + Catching Server Errors ~~~~~~~~~~~~~~~~~~~~~~ diff --git a/docs/guides/index.rst b/docs/guides/index.rst index d10172c..e90c9c5 100644 --- a/docs/guides/index.rst +++ b/docs/guides/index.rst @@ -10,4 +10,5 @@ and deployment. testing error-handling + oauth-security deployment diff --git a/docs/guides/oauth-security.rst b/docs/guides/oauth-security.rst new file mode 100644 index 0000000..07c5857 --- /dev/null +++ b/docs/guides/oauth-security.rst @@ -0,0 +1,319 @@ +OAuth Security and Deployment +============================== + +Deploying OAuth with the mcp-cpp-sdk requires careful attention to token management, network security, and issuer binding. This guide covers the operational controls and design rationale that developers and operators must understand to deploy OAuth safely. + +Token Persistence and Storage +------------------------------ + +The :cpp:class:`mcp::auth::InMemoryTokenStore` is thread-safe but volatile: tokens are lost when the process exits. Production deployments **must not** rely on it. + +Implementing a Persistent TokenStore +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Implement :cpp:class:`mcp::auth::TokenStore` with encryption and OS-appropriate secret storage: + +.. code-block:: cpp + + class MyTokenStore : public mcp::auth::TokenStore { + public: + void store(const std::string& server_url, + mcp::auth::TokenResponse token) override { + // Encrypt token and persist to secure storage + // (e.g., OS keychain, encrypted file, secret manager) + auto encrypted = encrypt_token(token); + persist_to_secure_storage(server_url, encrypted); + } + + std::optional load( + const std::string& server_url) const override { + // Retrieve and decrypt from secure storage + auto encrypted = retrieve_from_secure_storage(server_url); + if (!encrypted) return std::nullopt; + return decrypt_token(*encrypted); + } + + void remove(const std::string& server_url) override { + // Delete from secure storage + delete_from_secure_storage(server_url); + } + }; + +Never log access tokens, refresh tokens, authorization codes, client secrets, or PKCE verifiers. + +Credential-to-Issuer Binding +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The :cpp:class:`mcp::auth::ClientCredentialStore` is **deliberately separate** from :cpp:class:`mcp::auth::TokenStore` because they use different storage keys: + +- **TokenStore** is keyed by MCP server URL. +- **ClientCredentialStore** is keyed by authorization-server issuer. + +This separation (SEP-2352) prevents a vulnerability where a single MCP server protected by different authorization servers over time could accidentally present one server's credentials to another. One authorization server can also protect multiple MCP server URLs, so server-URL keying would eventually misattribute credentials. + +**Do not merge these stores.** Use separate persistent implementations, keyed correctly, and validate that a stored credential's issuer matches the issuer being contacted before presenting the credential. The SDK's :cpp:func:`mcp::auth::select_client_identity` function validates the issuer binding automatically when reusing stored credentials; operators who implement custom storage must enforce the same check. + +The Class That Applies These Controls +-------------------------------------- + +Everything above describes controls rather than the code that enforces them. :cpp:class:`mcp::auth::OAuthAuthorizationManager` is that code, and the only supported way to act on a ``WWW-Authenticate`` challenge. It applies the fetch policy to every request it makes, binds the issuer byte-for-byte to the metadata document it was read from, generates and checks a cryptographic ``state``, applies RFC 9207 ``iss`` validation before any ``error`` field in the response is read, carries S256 PKCE, and sends the RFC 8707 ``resource`` indicator into the code exchange. + +Composing :cpp:class:`mcp::auth::OAuthDiscoveryClient` and :cpp:func:`mcp::auth::OAuthHttpClient::exchange_code` into your own flow compiles and appears to work while applying **none** of them. That is the failure mode this guide exists to prevent: the fetch policy alone governs which targets are contacted, and says nothing about whether what comes back is bound to the issuer that was asked for. + +Worked Flow +~~~~~~~~~~~ + +Taken from ``examples/features/oauth_flow.cpp``, which runs end to end against a mock authorization server: + +.. literalinclude:: ../../examples/features/oauth_flow.cpp + :language: cpp + :start-after: docs-begin: manager-flow + :end-before: docs-end: manager-flow + :dedent: 8 + +The consent callback in that example mints the authorization code directly from its own mock server, which is possible only because the example owns both ends. A real client opens ``request.authorization_url`` in the user's browser, waits for the redirect to ``request.redirect_uri``, and returns :cpp:func:`mcp::auth::parse_authorization_response`. It must never invent ``state`` or ``iss``: the manager compares both against what it recorded and rejects a response that does not match. + +After a successful challenge, :cpp:func:`mcp::auth::OAuthAuthorizationManager::last_authorization_request` returns the record the attempt was validated against, so an application can audit the binding rather than take it on trust. + +See :doc:`/concepts/oauth` for the full API reference and for what the low-level building blocks do and do not provide. + +Metadata-Fetch Policy and Origin Validation +-------------------------------------------- + +Every OAuth metadata fetch — challenge-supplied `resource_metadata` URLs, protected-resource discovery, authorization-server discovery, and token-endpoint requests derived from discovered metadata — is an outbound-request primitive controlled by the :cpp:struct:`mcp::auth::MetadataFetchPolicy`. + +Policy Default Behavior +~~~~~~~~~~~~~~~~~~~~~~~~ + +A default-constructed policy **denies every origin**: :cpp:member:`mcp::auth::MetadataFetchPolicy::allowed_origins` is empty, and an origin not explicitly listed is refused **before host resolution**. When the URL is refused, no socket is opened and no DNS lookup occurs. + +.. code-block:: cpp + + mcp::auth::MetadataFetchPolicy policy; + // Every origin is denied; no metadata fetch succeeds without configuration. + +Configuring Allowed Origins +~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +List the exact origins your deployment intends to contact: + +.. code-block:: cpp + + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = { + "https://accounts.example.com", + "https://auth.internal:8443", + "https://protected-resource.example.org" + }; + +Comparison canonicalizes scheme and host case, a trailing run of host dots, a strictly-numeric explicit port that equals the scheme's default (``:443`` for ``https``, ``:80`` for ``http`` — so ``:00443`` and ``:0443`` both canonicalize the same as no port at all), and an IP literal's textual form (an expanded and a compressed IPv6 spelling of the same address compare equal). A non-default port and a genuinely different host still distinguish origins exactly. An allow-list entry that carries a path, query or fragment after the authority (e.g. ``"https://as.test/realms/foo"``) is not a bare origin and matches nothing. + +Deny List Always Wins +~~~~~~~~~~~~~~~~~~~~~ + +The deny list is consulted **before** the allow list. An origin in both lists is refused: + +.. code-block:: cpp + + policy.allowed_origins = {"https://auth.example.com"}; + policy.denied_origins = {"https://auth.example.com"}; + // Result: origin is DENIED + // Use this to temporarily block a server without removing it from the allow list. + +Runtime Override with origin_allowance Callback +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Applications that cannot enumerate all authorization servers in advance can implement the :cpp:member:`mcp::auth::MetadataFetchPolicy::origin_allowance` callback. It is consulted only when an origin is not in the allow list, and its return value determines whether the origin is allowed: + +.. code-block:: cpp + + policy.origin_allowance = [](const std::string& origin) { + // Custom rule: allow any HTTPS origin from internal domain. The origin here is already + // canonicalized, so a default port (e.g. "https://auth.internal:443") has been elided down + // to "https://auth.internal" — match on ".internal", not ".internal:", so a default-port + // origin is not missed. + if (origin.find("https://") == 0 && + origin.find(".internal") != std::string::npos) { + return true; + } + // Otherwise deny + return false; + }; + +The callback: + +- **Cannot relax scheme validation**: it cannot permit plain `http://` when the policy enforces HTTPS. +- **Cannot relax address validation**: it cannot permit loopback, private ranges, link-local, multicast, or reserved addresses after they have been rejected by address classification. +- Runs **after** the deny list, so denied origins stay denied. +- Runs **before** host resolution, so a rejected origin never results in a socket opening. + +Metadata URLs named by a protected resource are attacker-influenced input; the callback implements the rule that would otherwise be written as a complete allow list. Do not use it as a blanket escape hatch. + +Issuer Comparison and Canonical Form +------------------------------------- + +When an authorization server advertises support for RFC 9207 `iss`, the SDK validates that the issuer in the authorization response matches the issuer recorded from the selected authorization server's metadata. This comparison is **a plain byte-for-byte string equality check**: no canonicalization is performed. + +Why No Canonicalization +~~~~~~~~~~~~~~~~~~~~~~~~ + +String equality without URL normalization might seem unsafe. Consider this threat: + +1. Attacker registers `https://attacker.evil` as an authorization server in operator's metadata-fetch policy. +2. Attacker directs victim to `https://attacker-evil.com` (note the difference). +3. The two domains resolve to the same IP address or are aliased via DNS. +4. If the SDK normalized URLs or resolved hostnames to compare, `https://attacker.evil` and `https://attacker-evil.com` might appear equivalent and bypass the issuer check. +5. Attacker's server could then issue tokens claiming to be from a trusted issuer, or replace legitimate issuer URLs with attacker-controlled variants. + +By comparing only the byte sequence as written, the SDK ensures that: + +- An issuer URL is only accepted if it matches **exactly** what was advertised in the trusted metadata document. +- Hostname canonicalization, default-port elision, trailing-slash normalization, or percent-encoding variations cannot defeat the binding. +- No DNS-like resolution step can make two different issuer URLs appear equivalent. + +The issuer binding is only as strong as the metadata document it comes from. The SDK validates the metadata URL itself (origin, address, redirect path) before fetching; that separate validation is your defense against SSRF attacks on the metadata endpoint. + +Localhost Plain-HTTP Loopback Opt-Out +-------------------------------------- + +The :cpp:member:`mcp::auth::MetadataFetchPolicy::allow_plain_http_loopback` flag permits plain `http://` URLs and loopback addresses. This is a narrow opt-out intended **only for loopback development and test fixtures** such as the official MCP conformance runner. + +.. code-block:: cpp + + // Development / test fixture only + mcp::auth::MetadataFetchPolicy policy; + policy.allow_plain_http_loopback = true; // Narrow opt-out + +**Production deployments must NOT enable this flag.** HTTPS is the default and the security boundary for OAuth. + +Once enabled: + +- Plain `http://` schemes are accepted for loopback addresses (127.0.0.1, ::1). +- Loopback addresses themselves are unblocked (they are otherwise refused as private ranges). +- This flag does **not** relax any other control: link-local, private non-loopback, multicast, reserved, or address-based denials still apply. Redirect chains, response sizes, and origin checks remain enforced. +- It is never enabled implicitly; an operator must explicitly set it to `true`. + +If you are integrating with the MCP conformance suite or a local test server, enable this flag only in the test or development configuration, never in production deployments. + +SSRF and DNS-Rebinding Prevention +---------------------------------- + +Metadata URLs are attacker-influenced input when they come from a challenge or a protected resource's `authorization_servers` list. The :cpp:struct:`mcp::auth::MetadataFetchPolicy` includes controls to prevent Server-Side Request Forgery (SSRF) and DNS-rebinding attacks. + +Blocked Address Classes +~~~~~~~~~~~~~~~~~~~~~~~ + +Before connecting to a resolved address, the SDK classifies it and refuses several ranges: + +- **Link-local** (169.254.0.0/16 for IPv4, fe80::/10 for IPv6): typically used for local-network autoconfiguration, not safe routing targets. +- **RFC 1918 private ranges** (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16 for IPv4) and IPv6 unique-local (fc00::/7): internal network addresses. +- **Loopback** (127.0.0.0/8 for IPv4, ::1 for IPv6): unless explicitly permitted by `allow_plain_http_loopback`. +- **Multicast and broadcast**: not unicast-routable. +- **Reserved, unspecified, or non-routable**: e.g., 0.0.0.0, 255.255.255.255, documentation prefixes. + +Addresses written as IP literals in the URL are classified without any DNS lookup. A URL naming `169.254.169.254` or a private-range IP is refused immediately. + +Resolve-Then-Pin +~~~~~~~~~~~~~~~~ + +When a hostname must be resolved, the SDK performs a single lookup, classifies each resolved address, and connects only to addresses that pass the check. Critically, **addresses obtained from that single resolution are pinned for the connection**: a second resolution of the same hostname cannot redirect an established fetch, preventing DNS-rebinding attacks where: + +1. First lookup of `attacker.com` returns a public IP. +2. Metadata fetch begins to that IP. +3. Second lookup of `attacker.com` returns a private IP. +4. Attacker tries to pivot the connection to the private IP. + +This cannot happen: the first resolution's results are used for the entire fetch. + +Bounded and Re-Validated Redirects +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +HTTP redirects are followed up to a configurable limit (default 3): + +.. code-block:: cpp + + policy.max_redirects = 3; // Default; adjust as needed + +Each redirect target URL is validated afresh against the scheme, origin, and address controls. A redirect chain cannot escape the policy by landing on a forbidden origin or address. + +Response Size Caps +~~~~~~~~~~~~~~~~~~ + +Metadata and token responses are capped to prevent unbounded buffering: + +.. code-block:: cpp + + policy.max_response_bytes = 256 * 1024; // Default: 256 KiB + +A response exceeding the cap is rejected without being fully buffered, protecting against slowloris or resource-exhaustion attacks. + +Example: Hardened Policy Configuration +--------------------------------------- + +A production authorization-server integration might configure: + +.. code-block:: cpp + + mcp::auth::MetadataFetchPolicy policy; + + // Explicit allow list of trusted authorization servers + policy.allowed_origins = { + "https://accounts.example.com", + "https://auth-backup.example.com" + }; + + // Temporary block during incident + policy.denied_origins = {"https://accounts.example.com:8080"}; + + // For discovered protected-resource metadata that names + // authorization servers not in the allow list, require explicit + // operator approval via a callback (optional; safer to omit) + policy.origin_allowance = [](const std::string& origin) { + // Only allow explicit HTTPS from trusted domain + return origin.find("https://trusted-domain.internal") == 0; + }; + + // Response limits + policy.max_response_bytes = 256 * 1024; + policy.max_redirects = 3; + + // Never set allow_plain_http_loopback in production + policy.allow_plain_http_loopback = false; // Explicitly documented + +Give this policy to the authorization manager, which applies it to every +discovery and token request it makes: + +.. code-block:: cpp + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = "https://mcp.example.com/mcp"; + config.client_id = "my-mcp-client"; + config.redirect_uri = "https://app.example.com/callback"; + config.policy = policy; + + auto manager = std::make_shared( + executor, token_store, std::move(config), authorize); + +An application driving an :cpp:class:`mcp::auth::OAuthHttpClient` directly must +install the policy on it, or every fetch is refused: + +.. code-block:: cpp + + auto http_client = std::make_shared(executor); + http_client->set_metadata_policy(policy); + +The fetch policy is only one of the controls an authorization flow needs. It +stops the SDK contacting a target it should not, but it says nothing about +whether the metadata that comes back is bound to the issuer that was asked for, +whether the ``state`` echoed by the authorization server matches, or whether the +RFC 9207 ``iss`` parameter is correct. +:cpp:class:`mcp::auth::OAuthAuthorizationManager` applies those; an application +composing the low-level pieces itself does not get them. See +:ref:`oauth-what-the-manager-validates`. + +Cross-References +---------------- + +- See :doc:`/concepts/oauth` for the challenge-driven authorization flow and API reference. +- See :doc:`/concepts/transports` for general HTTP transport security. +- See :doc:`deployment` for overall production deployment patterns. diff --git a/docs/guides/testing.rst b/docs/guides/testing.rst index ef84bfd..be5f87a 100644 --- a/docs/guides/testing.rst +++ b/docs/guides/testing.rst @@ -78,9 +78,17 @@ When writing scripted tests, consider the following: auto init_res = co_await client->connect(client_info, client_caps); EXPECT_EQ(init_res.serverInfo.version, "1.0"); - // 2. Call a tool and verify the content - auto result = co_await client->call_tool("echo", {{"message", "test"}}); - EXPECT_EQ(result.content[0]["text"], "test"); + // 2. Call a tool and verify the content. + // Build the arguments into a named variable: an initializer list inside + // a co_await expression crashes GCC 12 and 13. See "Compiler Notes" in + // the README. + nlohmann::json echo_args{{"message", "test"}}; + auto result = co_await client->call_tool("echo", echo_args); + + // CallToolResult::content holds ContentBlock, a variant, not JSON. + const auto* text = std::get_if(&result.content.at(0)); + ASSERT_NE(text, nullptr); + EXPECT_EQ(text->text, "test"); // 3. Verify server-side state if accessible EXPECT_TRUE(server->has_tool("echo")); diff --git a/docs/index.rst b/docs/index.rst index b6492f0..2c11f18 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -16,8 +16,9 @@ Features * **Flexible transports**: stdio, WebSocket, Streamable HTTP, and custom transport support * **Type-safe**: Strong typing with JSON serialization via nlohmann_json * **Async-first**: Built on Boost.Asio for high-performance I/O -* **Full MCP support**: Tools, resources, prompts, sampling, roots, progress, and notifications -* **Authentication-ready**: OAuth 2.1 helpers and authenticated client transport support +* **Broad MCP surface**: Tools, resources, prompts, sampling, roots, progress, and notifications +* **Measured conformance**: Pinned official suites with explicit expected-failure evidence +* **Authentication building blocks**: Experimental OAuth helpers and HTTP bearer validation Quick Links ----------- diff --git a/examples/clients/stdio/client_stdio.cpp b/examples/clients/stdio/client_stdio.cpp index e6cbc13..15c53ce 100644 --- a/examples/clients/stdio/client_stdio.cpp +++ b/examples/clients/stdio/client_stdio.cpp @@ -1,6 +1,11 @@ /// @file client_stdio.cpp /// @brief MCP client over stdio demonstrating connect, list/call tools, /// list/read resources, list templates, list/get prompts, complete, and notifications. +/// +/// @note On the stdio transport, stdout IS the protocol channel: StdioTransport +/// writes JSON-RPC messages there by default. Every human-readable line in +/// this file therefore goes to stderr. Diagnostics printed to stdout would +/// interleave prose with JSON-RPC and corrupt the stream for the peer. #include #include @@ -44,14 +49,14 @@ int main() { // NOLINT(readability-function-cognitive-complexity) caps.sampling = ClientCapabilities::SamplingCapability{}; auto init_result = co_await client.connect(client_info, caps); - std::cout << "Connected to: " << init_result.serverInfo.name << " " + std::cerr << "Connected to: " << init_result.serverInfo.name << " " << init_result.serverInfo.version << "\n"; // list_tools auto tools = co_await client.list_tools(); - std::cout << "Tools (" << tools.tools.size() << "):\n"; + std::cerr << "Tools (" << tools.tools.size() << "):\n"; for (const auto& tool : tools.tools) { - std::cout << " - " << tool.name << "\n"; + std::cerr << " - " << tool.name << "\n"; } // call_tool with typed arguments @@ -59,7 +64,7 @@ int main() { // NOLINT(readability-function-cognitive-complexity) co_await client.call_tool("add", AddArgs{.a = 3.0, .b = 4.0} /* NOLINT */); for (const auto& block : add_result.content) { if (const auto* text = std::get_if(&block)) { - std::cout << "add(3, 4) = " << text->text << "\n"; + std::cerr << "add(3, 4) = " << text->text << "\n"; } } @@ -68,15 +73,15 @@ int main() { // NOLINT(readability-function-cognitive-complexity) auto echo_result = co_await client.call_tool("echo_async", echo_args); for (const auto& block : echo_result.content) { if (const auto* text = std::get_if(&block)) { - std::cout << "echo_async: " << text->text << "\n"; + std::cerr << "echo_async: " << text->text << "\n"; } } // list_resources auto resources = co_await client.list_resources(); - std::cout << "Resources (" << resources.resources.size() << "):\n"; + std::cerr << "Resources (" << resources.resources.size() << "):\n"; for (const auto& res : resources.resources) { - std::cout << " - " << res.uri << " (" << res.name << ")\n"; + std::cerr << " - " << res.uri << " (" << res.name << ")\n"; } // read_resource @@ -84,23 +89,23 @@ int main() { // NOLINT(readability-function-cognitive-complexity) auto read_result = co_await client.read_resource(resources.resources[0].uri); for (const auto& content : read_result.contents) { if (const auto* text = std::get_if(&content)) { - std::cout << "Resource content: " << text->text << "\n"; + std::cerr << "Resource content: " << text->text << "\n"; } } } // list_resource_templates auto templates = co_await client.list_resource_templates(); - std::cout << "Resource templates (" << templates.resourceTemplates.size() << "):\n"; + std::cerr << "Resource templates (" << templates.resourceTemplates.size() << "):\n"; for (const auto& tmpl : templates.resourceTemplates) { - std::cout << " - " << tmpl.uriTemplate << " (" << tmpl.name << ")\n"; + std::cerr << " - " << tmpl.uriTemplate << " (" << tmpl.name << ")\n"; } // list_prompts auto prompts = co_await client.list_prompts(); - std::cout << "Prompts (" << prompts.prompts.size() << "):\n"; + std::cerr << "Prompts (" << prompts.prompts.size() << "):\n"; for (const auto& prompt : prompts.prompts) { - std::cout << " - " << prompt.name << "\n"; + std::cerr << " - " << prompt.name << "\n"; } // get_prompt @@ -113,7 +118,7 @@ int main() { // NOLINT(readability-function-cognitive-complexity) co_await client.get_prompt(prompts.prompts[0].name, std::move(prompt_args)); for (const auto& msg : prompt_result.messages) { if (const auto* text = std::get_if(&msg.content)) { - std::cout << "Prompt message: " << text->text << "\n"; + std::cerr << "Prompt message: " << text->text << "\n"; } } } @@ -130,11 +135,11 @@ int main() { // NOLINT(readability-function-cognitive-complexity) complete_params.ref = std::move(prompt_ref); complete_params.argument = std::move(complete_arg); auto complete_result = co_await client.complete(complete_params); - std::cout << "Completions: " << complete_result.completion.values.size() << " values\n"; + std::cerr << "Completions: " << complete_result.completion.values.size() << " values\n"; // send_notification co_await client.send_notification("notifications/cancelled", std::nullopt); - std::cout << "Sent cancellation notification\n"; + std::cerr << "Sent cancellation notification\n"; }, boost::asio::detached); diff --git a/examples/clients/stdio/interactive_client.cpp b/examples/clients/stdio/interactive_client.cpp index e5489b4..e4e868b 100644 --- a/examples/clients/stdio/interactive_client.cpp +++ b/examples/clients/stdio/interactive_client.cpp @@ -1,3 +1,11 @@ +/// @file interactive_client.cpp +/// @brief Minimal interactive MCP client over stdio. +/// +/// @note On the stdio transport, stdout IS the protocol channel: StdioTransport +/// writes JSON-RPC messages there by default. Every human-readable line in +/// this file therefore goes to stderr. Diagnostics printed to stdout would +/// interleave prose with JSON-RPC and corrupt the stream for the peer. + #include #include #include @@ -11,7 +19,7 @@ using namespace mcp; Task run_client(const std::string& server_path) { auto executor = co_await boost::asio::this_coro::executor; - std::cout << "Connecting to server: " << server_path << "..." << '\n'; + std::cerr << "Connecting to server: " << server_path << "..." << '\n'; // Simple stdio transport for demonstration. // In a real client, you would spawn the server process and connect its pipes. @@ -26,21 +34,21 @@ Task run_client(const std::string& server_path) { co_return; } - std::cout << "Connected! Listing tools..." << '\n'; + std::cerr << "Connected! Listing tools..." << '\n'; auto tools_result = co_await client.list_tools(); for (const auto& tool : tools_result.tools) { - std::cout << "- " << tool.name << ": " << tool.description.value_or("(no description)") << '\n'; + std::cerr << "- " << tool.name << ": " << tool.description.value_or("(no description)") << '\n'; } if (!tools_result.tools.empty()) { std::string tool_name = tools_result.tools[0].name; - std::cout << "Calling first tool: " << tool_name << "..." << '\n'; + std::cerr << "Calling first tool: " << tool_name << "..." << '\n'; nlohmann::json args = nlohmann::json::object(); try { auto result = co_await client.call_tool(tool_name, args); - std::cout << "Result: " << nlohmann::json(result.content).dump(2) << '\n'; + std::cerr << "Result: " << nlohmann::json(result.content).dump(2) << '\n'; } catch (const std::exception& e) { std::cerr << "Error calling tool: " << e.what() << '\n'; } diff --git a/examples/features/README.md b/examples/features/README.md index 4ef6289..a850caf 100644 --- a/examples/features/README.md +++ b/examples/features/README.md @@ -231,18 +231,22 @@ done ### 12. OAuth Flow (`oauth_flow.cpp`) **What it demonstrates:** -- Full OAuth flow with mock auth server -- `OAuthAuthenticator` with `InMemoryTokenStore` +- `OAuthAuthorizationManager` acting on a `WWW-Authenticate` challenge from a mock authorization server: discovery under a `MetadataFetchPolicy`, issuer binding, cryptographic `state`, S256 PKCE, RFC 9207 response validation, and an RFC 8707 resource-indicated code exchange +- An application-supplied consent callback standing in for a browser +- `InMemoryTokenStore` holding the acquired token - `OAuthClientTransport` for token injection - Token refresh on auth failures - `make_auth_middleware()` for server-side validation +- The server-side challenge API that produces what the client acts on: one `ProtectedResourceMetadataConfig` drives the metadata route via `protected_resource_metadata_path()`, the document body via `format_protected_resource_metadata()`, and the advertised URL via `protected_resource_metadata_url()`, while `format_www_authenticate()` renders the `BearerChallengeConfig` into the 401 header **Key APIs:** -- `OAuthAuthenticator` -- `InMemoryTokenStore` +- `OAuthAuthorizationManager` +- `AuthorizationCallback` +- `MetadataFetchPolicy` - `OAuthClientTransport` - `make_auth_middleware()` -- `generate_pkce_pair()` +- `BearerChallengeConfig` / `format_www_authenticate()` +- `ProtectedResourceMetadataConfig` / `protected_resource_metadata_path()` / `protected_resource_metadata_url()` / `format_protected_resource_metadata()` **Build:** `python scripts/build.py --examples` @@ -250,6 +254,8 @@ done **Expected output:** OAuth discovery, token acquisition, injection, refresh on expiry +**Serving the challenge over HTTP:** this example runs over `MemoryTransport`, so it renders the 401 header itself. `StreamableHttpSessionManager` and `HttpServerTransport` send it for you — call `set_protected_resource_metadata()`, `set_bearer_token_validator()` or `set_async_bearer_token_validator()`, and `set_unauthenticated_paths()` before `listen()`. See the "Protecting a server" section of `docs/concepts/oauth.rst`. + --- ## Common Patterns @@ -284,16 +290,16 @@ Examples such as ``http_server_convenience.cpp`` and ``graceful_shutdown.cpp`` d Async tool handlers with Context: ```cpp -server.add_tool("my_tool", "Description", schema, - [](mcp::Context& ctx, nlohmann::json args) -> mcp::Task { +server.add_tool("my_tool", "Description", schema, + [](mcp::Context& ctx, nlohmann::json args) -> mcp::Task { co_await ctx.log_info("Starting work..."); if (ctx.is_cancelled()) { - co_return make_error_result("Cancelled"); + co_return mcp::make_tool_error_result("Cancelled"); } co_await ctx.report_progress(50, 100); - co_return make_success_result("Done"); + co_return mcp::make_tool_text_result("Done"); }); ``` diff --git a/examples/features/error_handling.cpp b/examples/features/error_handling.cpp index c95dec8..d3d3735 100644 --- a/examples/features/error_handling.cpp +++ b/examples/features/error_handling.cpp @@ -59,17 +59,18 @@ int main() { {"properties", {{"code", {{"type", "integer"}}}}}, {"required", nlohmann::json::array({"code"})}}; - server.add_tool("return_error", "Returns error JSON", return_error_schema, - [](const nlohmann::json& args) -> nlohmann::json { - std::cout << "[Server] return_error handler called\n"; - int code = args.at("code").get(); - mcp::CallToolResult result; - mcp::TextContent content; - content.text = "Application error with code " + std::to_string(code); - result.content.emplace_back(std::move(content)); - result.isError = true; - return nlohmann::json(result); - }); + server.add_tool( + "return_error", "Returns error JSON", return_error_schema, + [](const nlohmann::json& args) -> mcp::CallToolResult { + std::cout << "[Server] return_error handler called\n"; + int code = args.at("code").get(); + mcp::CallToolResult result; + mcp::TextContent content; + content.text = "Application error with code " + std::to_string(code); + result.content.emplace_back(std::move(content)); + result.isError = true; + return result; + }); // Tool 3: Accesses invalid argument (nlohmann::json throws) nlohmann::json invalid_access_schema = { @@ -77,17 +78,17 @@ int main() { {"properties", {{"optional_field", {{"type", "string"}}}}}, {"required", nlohmann::json::array({})}}; - server.add_tool("invalid_access", "Accesses non-existent argument", invalid_access_schema, - [](const nlohmann::json& args) -> nlohmann::json { - std::cout << "[Server] invalid_access handler called\n"; - // This will throw if "required_field" doesn't exist - std::string value = args.at("required_field").get(); - mcp::CallToolResult result; - mcp::TextContent content; - content.text = value; - result.content.emplace_back(std::move(content)); - return nlohmann::json(result); - }); + server.add_tool( + "invalid_access", "Accesses non-existent argument", invalid_access_schema, + [](const nlohmann::json& args) -> mcp::CallToolResult { + std::cout << "[Server] invalid_access handler called\n"; + std::string value = args.at("required_field").get(); + mcp::CallToolResult result; + mcp::TextContent content; + content.text = value; + result.content.emplace_back(std::move(content)); + return result; + }); // Tool 4: Successful tool (for comparison) nlohmann::json success_schema = {{"type", "object"}, diff --git a/examples/features/graceful_shutdown.cpp b/examples/features/graceful_shutdown.cpp index 18701ab..d49bc68 100644 --- a/examples/features/graceful_shutdown.cpp +++ b/examples/features/graceful_shutdown.cpp @@ -87,7 +87,7 @@ int main() { server.add_tool("quick_task", "A quick task that completes immediately", quick_schema, [](const nlohmann::json& args) -> nlohmann::json { - std::cout << "[Server] quick_task called\n"; + std::cerr << "[Server] quick_task called\n"; int value = args.at("value").get(); CallToolResult result; TextContent content; @@ -104,7 +104,7 @@ int main() { server.add_tool( "slow_task", "A task that takes time to complete", slow_schema, [](nlohmann::json args) -> Task { - std::cout << "[Server] slow_task called\n"; + std::cerr << "[Server] slow_task called\n"; int duration_ms = args.at("duration_ms").get(); // Simulate work with async delay @@ -113,7 +113,7 @@ int main() { work_timer.expires_after(std::chrono::milliseconds(duration_ms)); co_await work_timer.async_wait(asio::use_awaitable); - std::cout << "[Server] slow_task completed after " << duration_ms << "ms\n"; + std::cerr << "[Server] slow_task completed after " << duration_ms << "ms\n"; CallToolResult result; TextContent content; content.text = "Slow task completed after " + std::to_string(duration_ms) + "ms"; @@ -122,6 +122,8 @@ int main() { }); // ========== TRANSPORT SETUP ========== + // stdout is handed to the transport, so it is the protocol channel and + // carries JSON-RPC only. Every diagnostic in this example goes to stderr. BlockingInputBuffer input_buffer; std::istream controlled_input(&input_buffer); auto transport = @@ -134,10 +136,10 @@ int main() { asio::steady_timer trigger(io_ctx.get_executor()); trigger.expires_after(std::chrono::seconds(1)); co_await trigger.async_wait(asio::use_awaitable); - std::cout << "[Main] Initiating graceful shutdown\n"; + std::cerr << "[Main] Initiating graceful shutdown\n"; input_buffer.close(); graceful_shutdown(io_ctx, transport); - std::cout << "[Main] Graceful shutdown initiated (timeout: " + std::cerr << "[Main] Graceful shutdown initiated (timeout: " << constants::g_shutdown_timeout_ms << "ms)\n"; }, asio::detached); @@ -146,21 +148,21 @@ int main() { asio::co_spawn( io_ctx, [&, transport]() -> Task { - std::cout << "[Server] Starting\n"; + std::cerr << "[Server] Starting\n"; try { co_await server.run(transport, io_ctx.get_executor()); - std::cout << "[Server] Stopped normally\n"; + std::cerr << "[Server] Stopped normally\n"; } catch (const std::exception& e) { - std::cout << "[Server] Stopped with exception: " << e.what() << "\n"; + std::cerr << "[Server] Stopped with exception: " << e.what() << "\n"; } }, asio::detached); // ========== RUN IO CONTEXT ========== - std::cout << "[Main] Starting io_context\n"; + std::cerr << "[Main] Starting io_context\n"; io_ctx.run(); - std::cout << "[Main] io_context stopped, exiting normally\n"; + std::cerr << "[Main] io_context stopped, exiting normally\n"; return EXIT_SUCCESS; } catch (const std::exception& e) { diff --git a/examples/features/middleware.cpp b/examples/features/middleware.cpp index 2acad55..ee21d0a 100644 --- a/examples/features/middleware.cpp +++ b/examples/features/middleware.cpp @@ -21,6 +21,7 @@ #include #include #include +#include #include namespace asio = boost::asio; @@ -83,13 +84,8 @@ int main() { params.at("arguments").at("block").get()) { std::cout << "[Middleware C] Short-circuiting: block flag detected\n"; co_await ctx.log_info("Middleware C: short-circuit"); - // Return error without calling next - mcp::CallToolResult err; - mcp::TextContent content; - content.text = "Request blocked by middleware C"; - err.content.emplace_back(std::move(content)); - err.isError = true; - co_return nlohmann::json(err); + // Exceptions are converted into protocol-valid tool errors. + throw std::runtime_error("Request blocked by middleware C"); } auto result = co_await next(ctx, params); @@ -108,12 +104,7 @@ int main() { server.add_tool("echo", "Echo a message", echo_schema, [](const nlohmann::json& args) -> nlohmann::json { std::cout << "[Handler] Executing echo tool\n"; - // Return a CallToolResult-compatible structure with content - mcp::CallToolResult result; - mcp::TextContent content; - content.text = args.at("message").get(); - result.content.emplace_back(std::move(content)); - return nlohmann::json(result); + return nlohmann::json{{"echo", args.at("message").get()}}; }); // ========== CLIENT SETUP ========== diff --git a/examples/features/oauth_flow.cpp b/examples/features/oauth_flow.cpp index ca9c817..477a6cb 100644 --- a/examples/features/oauth_flow.cpp +++ b/examples/features/oauth_flow.cpp @@ -3,17 +3,31 @@ /// /// This example shows: /// - In-process OAuth discovery + token endpoints served by a mock HTTP server -/// - OAuthDiscoveryClient resolving protected-resource and auth-server metadata -/// - OAuthAuthenticator with InMemoryTokenStore storing the exchanged token +/// - OAuthAuthorizationManager acting on a `WWW-Authenticate` challenge: discovery under a +/// MetadataFetchPolicy, issuer binding, cryptographic `state`, S256 PKCE, RFC 9207 response +/// validation, and an RFC 8707 resource-indicated code exchange +/// - An application-supplied consent callback standing in for a browser +/// - InMemoryTokenStore holding the acquired token, keyed by the MCP server URL /// - OAuthClientTransport wrapping a MemoryTransport client connection /// - Automatic auth-token injection into MCP requests /// - Automatic refresh after a server-side auth failure /// - make_auth_middleware() protecting MCP tools on the server side +/// - The server-side challenge API producing what the client acts on: a +/// ProtectedResourceMetadataConfig whose RFC 9728 3.1 path and URL come from +/// protected_resource_metadata_path() and protected_resource_metadata_url(), the document +/// itself from format_protected_resource_metadata(), and the `WWW-Authenticate` header from +/// a BearerChallengeConfig rendered by format_www_authenticate() +/// +/// @warning Authorization goes through OAuthAuthorizationManager. Driving OAuthDiscoveryClient +/// and OAuthHttpClient::exchange_code() by hand compiles and appears to work, but performs no +/// issuer binding, no `state` check and no RFC 9207 validation. See the comment in +/// run_client_demo() and the "Challenge-driven authorization" section of docs/concepts/oauth.rst. #include #include #include #include +#include #include #include @@ -121,12 +135,23 @@ class MockOAuthServer { port_ = acceptor_.local_endpoint().port(); issuer_ = "http://127.0.0.1:" + std::to_string(port_); protected_resource_url_ = issuer_ + "/memory-mcp"; + + // The one description of this protected resource. Both HTTP server transports take this + // same struct; here it drives the mock's metadata route and the challenge the client is + // handed, so the example exercises the derivation rather than restating its result. + resource_metadata_.resource = protected_resource_url_; + resource_metadata_.authorization_servers = {issuer_}; + resource_metadata_.scopes_supported = {"mcp:demo"}; } [[nodiscard]] const std::string& issuer() const { return issuer_; } [[nodiscard]] const std::string& protected_resource_url() const { return protected_resource_url_; } + [[nodiscard]] const mcp::ProtectedResourceMetadataConfig& resource_metadata() const { + return resource_metadata_; + } + [[nodiscard]] const std::vector& observed_tokens() const { return observed_tokens_; } [[nodiscard]] const std::string& last_observed_token() const { return last_observed_token_; } @@ -177,13 +202,25 @@ class MockOAuthServer { return response; } + static http::response make_json_body_response(http::status status, int version, + std::string body) { + http::response response{status, version}; + response.set(http::field::content_type, "application/json"); + response.keep_alive(false); + response.body() = std::move(body); + response.prepare_payload(); + return response; + } + http::response handle_request(const http::request& request) { + // RFC 9728 3.1 puts the well-known segment between the authority and the resource's own + // path, so a resource at /memory-mcp is described at + // /.well-known/oauth-protected-resource/memory-mcp. Deriving the route keeps the mock and + // the challenge from drifting apart, exactly as the server transports do. if (request.method() == http::verb::get && - request.target() == "/.well-known/oauth-protected-resource/memory-mcp") { - return make_json_response(http::status::ok, request.version(), - json{{"resource", protected_resource_url_}, - {"authorization_servers", json::array({issuer_})}, - {"scopes_supported", json::array({"mcp:demo"})}}); + request.target() == mcp::protected_resource_metadata_path(resource_metadata_)) { + return make_json_body_response(http::status::ok, request.version(), + mcp::format_protected_resource_metadata(resource_metadata_)); } if (request.method() == http::verb::get && @@ -279,6 +316,7 @@ class MockOAuthServer { unsigned short port_{}; std::string issuer_; std::string protected_resource_url_; + mcp::ProtectedResourceMetadataConfig resource_metadata_; std::map> authorization_codes_; std::vector observed_tokens_; std::string last_observed_token_; @@ -341,15 +379,10 @@ struct ClientFlowRuntime { }; struct ClientFlowState { - std::shared_ptr oauth_http; std::shared_ptr token_store; - std::shared_ptr authenticator; - std::optional protected_metadata; - std::optional auth_metadata; + std::shared_ptr manager; std::optional stored_after_refresh; - mcp::auth::OAuthConfig oauth_config; - mcp::auth::PkcePair pkce; - std::string auth_code; + std::string resource; }; std::string make_initialize_request_wire() { @@ -393,46 +426,95 @@ auto run_client_demo(ClientFlowRuntime runtime) -> mcp::Task { std::cout << "OAuth mock server: " << runtime.mock_oauth->issuer() << '\n'; std::cout << "MCP client/server transport: MemoryTransport loopback\n\n"; - state->oauth_http = - std::make_shared(runtime.io_ctx->get_executor()); - mcp::auth::OAuthDiscoveryClient discovery(state->oauth_http); - - state->protected_metadata = co_await discovery.discover_protected_resource( - runtime.mock_oauth->protected_resource_url()); - state->auth_metadata = co_await discovery.discover_auth_server( - state->protected_metadata->authorization_servers.front()); - - std::cout << "[Client] Discovered protected resource: " << state->protected_metadata->resource - << '\n'; - std::cout << "[Client] Discovered token endpoint: " << state->auth_metadata->token_endpoint - << '\n'; - - state->pkce = mcp::auth::generate_pkce_pair(); - state->auth_code = runtime.mock_oauth->issue_authorization_code(state->pkce.code_verifier); - std::cout << "[Client] Generated PKCE challenge: " << state->pkce.code_challenge << '\n'; - - state->oauth_config.client_id = "oauth-flow-example-client"; - state->oauth_config.token_endpoint = state->auth_metadata->token_endpoint; - state->oauth_config.authorization_endpoint = state->auth_metadata->authorization_endpoint; - state->oauth_config.redirect_uri = "http://localhost/callback"; - state->oauth_config.scope = "mcp:demo"; + state->resource = runtime.mock_oauth->protected_resource_url(); + + // --------------------------------------------------------------------------------- + // Authorization runs through OAuthAuthorizationManager. Do not hand-roll this flow: + // driving OAuthDiscoveryClient and OAuthHttpClient::exchange_code() yourself compiles + // and appears to work, but skips the checks the manager performs: + // + // * the metadata issuer is bound byte-for-byte to the URL the document was fetched from; + // * `state` comes from a cryptographic random source and must match the response; + // * RFC 9207 `iss` is validated before the response's `error` fields are read; + // * PKCE S256 is generated and the verifier is retained for the token request; + // * the RFC 8707 `resource` indicator is carried into the code exchange. + // --------------------------------------------------------------------------------- + + // docs-begin: manager-flow + // A default-constructed policy refuses every metadata target: the allow list is empty + // and plain HTTP to a loopback address is not permitted. The application must name the + // origins it intends to reach before the first request. This example talks only to its + // own mock server on loopback, so it allows that one origin and opts in to plain-HTTP + // loopback. A real client lists the origins of its MCP server and of the authorization + // servers that server names, and leaves allow_plain_http_loopback false so only https + // targets are reachable. + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {runtime.mock_oauth->issuer()}; + policy.allow_plain_http_loopback = true; + + mcp::auth::OAuthAuthorizationConfig auth_config; + auth_config.server_url = state->resource; + auth_config.client_id = "oauth-flow-example-client"; + auth_config.redirect_uri = "http://localhost/callback"; + auth_config.policy = std::move(policy); state->token_store = std::make_shared(); - state->authenticator = std::make_shared( - state->token_store, state->oauth_http, state->oauth_config, - state->protected_metadata->resource); - { - auto initial_token = co_await state->oauth_http->exchange_code( - state->oauth_config, state->auth_code, state->pkce.code_verifier); - state->authenticator->store_token(initial_token); + // The consent step. The SDK never launches a browser or binds a listener for the + // redirect; that is the application's job. + // + // DO NOT COPY THIS CALLBACK. It mints the code directly from the mock authorization + // server, which is only possible because this example owns both ends. A real client + // opens request.authorization_url, waits for the redirect to request.redirect_uri, and + // returns parse_authorization_response(redirect_url). `state` and `iss` are echoed here + // as a genuine authorization server would echo them; the manager rejects the response + // if they do not match what it recorded. + auto* mock_oauth = runtime.mock_oauth; + auto authorize = [mock_oauth](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + std::cout << "[Client] Consent step for recorded issuer: " << request.issuer << '\n'; + std::cout << "[Client] Resource indicator: " << request.resource.value_or("") << '\n'; + + mcp::auth::AuthorizationResponse response; + response.code = mock_oauth->issue_authorization_code(request.code_verifier); + response.state = request.state; + response.iss = request.issuer; + co_return response; + }; + + state->manager = std::make_shared( + runtime.io_ctx->get_executor(), state->token_store, std::move(auth_config), + std::move(authorize)); + + // What a conforming MCP resource server returns with its 401, per RFC 9728. Over an HTTP + // transport StreamableHttpSessionManager and HttpServerTransport send this header for you + // once set_protected_resource_metadata() and set_bearer_challenge() are configured; this + // example runs over MemoryTransport, so it renders the same header through the same public + // functions those transports call. + mcp::BearerChallengeConfig challenge; + challenge.scope = "mcp:demo"; + challenge.resource_metadata = + mcp::protected_resource_metadata_url(runtime.mock_oauth->resource_metadata()); + const std::string www_authenticate = mcp::format_www_authenticate(challenge); + + std::cout << "[Client] Acting on the resource server's WWW-Authenticate challenge\n"; + if (!co_await state->manager->try_handle_challenge(www_authenticate)) { + throw std::runtime_error("challenge carried no Bearer authorization to act on"); } - std::cout << "[Client] Stored initial access token: " - << state->authenticator->get_access_token() << "\n"; + // docs-end: manager-flow + + // The manager records what each attempt was actually validated against, so an application + // can audit the binding rather than take it on trust. + const auto record = state->manager->last_authorization_request(); + if (!record) { + throw std::runtime_error("authorization completed without recording a request"); + } + std::cout << "[Client] Authorization succeeded against issuer: " << record->issuer << '\n'; + std::cout << "[Client] Stored initial access token (value redacted)\n"; { auto oauth_transport = std::make_shared( - runtime.client_base_transport, state->authenticator); + runtime.client_base_transport, state->manager); co_await oauth_transport->write_message(make_initialize_request_wire()); @@ -459,7 +541,7 @@ auto run_client_demo(ClientFlowRuntime runtime) -> mcp::Task { throw std::runtime_error("secure_echo returned an unexpected tool error"); } - state->stored_after_refresh = state->token_store->load(state->protected_metadata->resource); + state->stored_after_refresh = state->token_store->load(state->resource); if (!state->stored_after_refresh || state->stored_after_refresh->access_token != "refreshed-access-token") { throw std::runtime_error("token refresh did not persist the refreshed access token"); @@ -484,17 +566,9 @@ auto run_client_demo(ClientFlowRuntime runtime) -> mcp::Task { oauth_transport->close(); } - std::cout << "[Client] Observed tokens on the server: "; - for (std::size_t i = 0; i < runtime.mock_oauth->observed_tokens().size(); ++i) { - if (i != 0) { - std::cout << " -> "; - } - std::cout << runtime.mock_oauth->observed_tokens()[i]; - } - std::cout << '\n'; - - std::cout << "[Client] Refreshed access token in store: " - << state->stored_after_refresh->access_token << '\n'; + std::cout << "[Client] Server observed " << runtime.mock_oauth->observed_tokens().size() + << " bearer-token attempts (values redacted)\n"; + std::cout << "[Client] Refreshed access token is present in the store (value redacted)\n"; std::cout << "[Client] OAuth flow completed successfully\n"; runtime.mock_oauth->stop(); @@ -535,8 +609,8 @@ int main() { server.use( mcp::auth::make_auth_middleware([&mock_oauth](const std::string& token) -> Task { const bool valid = mock_oauth.validate_token(token); - std::cout << "[Server] Validating token: " << token - << (valid ? " (accepted)" : " (rejected)") << '\n'; + std::cout << "[Server] Validating bearer token (value redacted): " + << (valid ? "accepted" : "rejected") << '\n'; co_return valid; })); @@ -545,11 +619,10 @@ int main() { json{{"type", "object"}, {"properties", {{"message", {{"type", "string"}}}}}, {"required", json::array({"message"})}}, - [&mock_oauth](const json& args) -> CallToolResult { + [](const json& args) -> CallToolResult { return make_text_result("Server handled '" + args.at("message").get() + - "' using token " + mock_oauth.last_observed_token(), - std::nullopt, - json{{"tokenSeen", mock_oauth.last_observed_token()}}); + "' with an authenticated request", + std::nullopt, json{{"authenticated", true}}); }); auto [server_base_transport, client_base_transport] = diff --git a/examples/servers/stdio/server_stdio.cpp b/examples/servers/stdio/server_stdio.cpp index 4926169..a7192f9 100644 --- a/examples/servers/stdio/server_stdio.cpp +++ b/examples/servers/stdio/server_stdio.cpp @@ -144,6 +144,39 @@ int main() { server.add_resource_template(user_template); + // docs-begin: untyped-prompt + // Prompt: greet (untyped JSON handler) + Prompt greet_prompt; + greet_prompt.name = "greet"; + greet_prompt.description = "Greet someone by name"; + + PromptArgument who_arg; + who_arg.name = "who"; + who_arg.description = "Name of the person to greet"; + who_arg.required = true; + + greet_prompt.arguments = std::vector{std::move(who_arg)}; + + server.add_prompt( + greet_prompt, [](const nlohmann::json& params) -> nlohmann::json { + // params is the whole GetPromptRequestParams, not just the arguments: + // {"name":"greet","arguments":{"who":"world"}} + const auto arguments = params.value("arguments", nlohmann::json::object()); + const auto who = arguments.value("who", std::string{"world"}); + + PromptMessage msg; + msg.role = Role::eUser; + TextContent text_content; + text_content.text = "Say hello to " + who + "."; + msg.content = std::move(text_content); + + GetPromptResult result; + result.description = "Greeting prompt"; + result.messages.push_back(std::move(msg)); + return result; + }); + // docs-end: untyped-prompt + // Prompt: code_review (parameterized) Prompt review_prompt; review_prompt.name = "code_review"; diff --git a/include/mcp/auth/challenge.hpp b/include/mcp/auth/challenge.hpp new file mode 100644 index 0000000..d5eb469 --- /dev/null +++ b/include/mcp/auth/challenge.hpp @@ -0,0 +1,164 @@ +#pragma once + +#include +#include + +#include +#include +#include +#include + +namespace mcp::auth { + +/** + * @brief A single authentication challenge parsed from a `WWW-Authenticate` header. + * + * @details Parameter names are matched case-insensitively and stored lower-cased. Quoted values are + * unescaped. The named accessors below expose the parameters the MCP authorization flow consumes; + * `parameters` retains every parameter in the order it appeared. + */ +struct BearerChallenge { + std::string scheme; ///< Authentication scheme exactly as written, e.g. `Bearer`. + std::optional realm; ///< Optional `realm` parameter. + std::optional resource_metadata; ///< Optional RFC 9728 `resource_metadata` URL. + std::optional scope; ///< Optional space-delimited `scope` parameter. + std::optional error; ///< Optional OAuth `error` code. + std::optional error_description; ///< Optional human-readable error description. + std::optional error_uri; ///< Optional URL describing the error. + KeyValuePairList parameters; ///< Every parameter, lower-cased names, source order preserved. + + /// @return True when the scheme is `Bearer`, compared case-insensitively. + [[nodiscard]] MCP_API bool is_bearer() const; +}; + +/** + * @brief Parse one `WWW-Authenticate` header value into its constituent challenges. + * + * @param header_value Raw header value. + * @return Every challenge found, in source order. + * + * @details Handles quoted values with backslash escapes, case-insensitive parameter names, + * arbitrary parameter ordering, and several comma-separated challenges in one header. Malformed + * trailing input is discarded rather than throwing, so a partially understood header still yields + * the challenges that precede the damage. + */ +[[nodiscard]] MCP_API std::vector parse_www_authenticate( + std::string_view header_value); + +/** + * @brief Parse several `WWW-Authenticate` header values into one ordered challenge list. + * + * @param header_values Raw header values, in the order the response carried them. + * @return Every challenge found across all headers, in source order. + */ +[[nodiscard]] MCP_API std::vector parse_www_authenticate( + const std::vector& header_values); + +/** + * @brief Select the challenge the MCP authorization flow should act on. + * + * @param challenges Challenges parsed from the response. + * @return The first `Bearer` challenge, or `std::nullopt` when none is present. + */ +[[nodiscard]] MCP_API std::optional select_bearer_challenge( + const std::vector& challenges); + +/** + * @brief Per-attempt record of an authorization request. + * + * @details One record is created per authorization attempt and holds every value the response must be + * validated against: the CSRF `state`, the PKCE code verifier, and the issuer recorded from the + * selected authorization server's metadata document. The recorded issuer must come from metadata that + * was itself fetched and validated. + */ +struct AuthorizationRequest { + std::string authorization_url; ///< Complete authorization endpoint URL including query. + std::string state; ///< CSRF state generated from a cryptographic random source. + std::string code_verifier; ///< PKCE verifier retained for the token request. + std::string code_challenge; ///< PKCE challenge sent to the authorization endpoint. + std::string issuer; ///< Issuer recorded from the selected authorization server. + /// Value of `authorization_response_iss_parameter_supported` in that same metadata document. + bool issuer_parameter_supported{false}; + std::string client_id; ///< Client identifier used for this attempt. + std::string redirect_uri; ///< Redirect URI registered for this attempt. + std::optional scope; ///< Requested scope, omitted when no scope was selected. + std::optional resource; ///< RFC 8707 resource indicator for this attempt. +}; + +/** + * @brief Authorization response returned to the application's redirect handler. + */ +struct AuthorizationResponse { + std::optional code; ///< Authorization code on success. + std::optional state; ///< State echoed by the authorization server. + std::optional iss; ///< RFC 9207 issuer identifier, when supplied. + std::optional error; ///< OAuth error code on failure. + std::optional error_description; ///< Human-readable error description. + std::optional error_uri; ///< URL describing the error. +}; + +/** + * @brief Parse an authorization response from a redirect URL or a bare query string. + * + * @param redirect_url Redirect URL the authorization server sent the user agent to, or the query + * component on its own. + * @return The decoded response parameters. + * + * @details Values are decoded as `application/x-www-form-urlencoded`, so `+` becomes a space and + * percent-escapes are expanded. Decoding happens before any comparison, exactly as RFC 9207 + * Section 2.4 requires. + */ +[[nodiscard]] MCP_API AuthorizationResponse +parse_authorization_response(const std::string& redirect_url); + +/** + * @brief Outcome of validating an authorization response against its request record. + */ +enum class AuthorizationResponseStatus { + accepted, ///< Response is authentic and carries a code. + state_missing, ///< Response omitted `state` although the request sent one. + state_mismatch, ///< Returned `state` did not match the recorded value. + issuer_missing, ///< Server advertised `iss` support but omitted the parameter. + issuer_mismatch, ///< Returned `iss` did not match the recorded issuer. + server_error, ///< Authentic response carrying an OAuth `error`. + code_missing, ///< Authentic success response with no authorization code. +}; + +/** + * @brief Result of validating an authorization response. + */ +struct AuthorizationResponseValidation { + AuthorizationResponseStatus status{AuthorizationResponseStatus::accepted}; ///< Outcome. + std::string message; ///< Diagnostic description of the outcome. + + /// @return True when the authorization code may be sent to a token endpoint. + [[nodiscard]] bool accepted() const { return status == AuthorizationResponseStatus::accepted; } +}; + +/** + * @brief Validate an authorization response against the record of the request that produced it. + * + * @param request Per-attempt record holding the expected state and recorded issuer. + * @param response Response returned to the application's redirect handler. + * @return The validation outcome; only `accepted` permits transmitting the code. + * + * @details Applies RFC 9207 Section 2.4 as adopted by MCP, before any other interpretation of the + * response: + * + * | `authorization_response_iss_parameter_supported` | `iss` present | Action | + * |--------------------------------------------------|---------------|----------------------------| + * | `true` | yes | Compare to recorded issuer | + * | `true` | no | Reject | + * | `false` or absent | yes | Compare to recorded issuer | + * | `false` or absent | no | Proceed | + * + * Comparison is byte-for-byte string equality with no normalization of any kind. + * + * The issuer check runs ahead of the `error` path: a response carrying `error`, `error_description` + * or `error_uri` with a mismatched issuer is rejected as `issuer_mismatch`, and the caller must + * neither act on nor display those values. + */ +[[nodiscard]] MCP_API AuthorizationResponseValidation validate_authorization_response( + const AuthorizationRequest& request, const AuthorizationResponse& response); + +} // namespace mcp::auth diff --git a/include/mcp/auth/client_identity.hpp b/include/mcp/auth/client_identity.hpp new file mode 100644 index 0000000..1b88c22 --- /dev/null +++ b/include/mcp/auth/client_identity.hpp @@ -0,0 +1,267 @@ +#pragma once + +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace mcp::auth { + +/** + * @brief How the client identity presented to an authorization server was obtained. + */ +enum class ClientIdentitySource { + pre_registered, ///< Supplied by the application; never re-registered. + client_id_metadata_document, ///< A published HTTPS document URL used as the client identifier. + dynamic_registration, ///< Obtained from the server's RFC 7591 registration endpoint. +}; + +/** + * @brief Short, stable description of a client identity source. + * + * @param source Source to describe. + * @return A description suitable for diagnostics. + */ +[[nodiscard]] MCP_API std::string_view describe(ClientIdentitySource source); + +/** + * @brief Client metadata sent at dynamic registration and published as a client ID metadata + * document. + * + * @details The defaults describe an authorization-code client that also intends to renew its grant: + * `refresh_token` is advertised in `grant_types` whether or not the server issues one. + */ +struct OAuthClientMetadata { + std::vector redirect_uris; ///< Redirect URIs the client will use. + /// SEP-837 `application_type`. Always present in the registration request body. + std::string application_type{"native"}; + /// Advertised grant types. `refresh_token` is advertised so the server may issue one. + std::vector grant_types{"authorization_code", "refresh_token"}; + std::vector response_types{"code"}; ///< Advertised response types. + std::optional client_name; ///< Optional human-readable client name. + std::optional client_uri; ///< Optional client home page. + std::optional software_id; ///< Optional stable software identifier. + std::optional software_version; ///< Optional software version. + std::optional scope; ///< Optional space-delimited requested scope. + /// Optional token endpoint authentication method the client wishes to register for. + std::optional token_endpoint_auth_method; +}; + +/** + * @brief Serialize client metadata into an RFC 7591 registration request body. + * + * @param json JSON object to populate. + * @param metadata Metadata to serialize. + */ +MCP_API void to_json(nlohmann::json& json, const OAuthClientMetadata& metadata); + +/** + * @brief Client credentials together with the issuer they belong to. + * + * @details SEP-2352 binds credentials to one authorization server. `issuer` records that binding so + * a credential can never be presented to a different server by accident, and so a stored entry + * whose recorded issuer disagrees with its storage key can be discarded rather than trusted. + */ +struct OAuthClientInformation { + std::string client_id; ///< Client identifier presented to the server. + std::optional client_secret; ///< Optional confidential-client secret. + std::optional client_id_issued_at; ///< Issue time reported by the server. + std::optional client_secret_expires_at; ///< Secret expiry, 0 meaning "never". + std::string issuer; ///< Issuer these credentials are bound to. + ClientIdentitySource source{ClientIdentitySource::dynamic_registration}; ///< How they were got. + + /** + * @brief Determine whether the server-reported secret expiry has passed. + * + * @param now_seconds Current time as seconds since the Unix epoch. + * @return True when a non-zero expiry was reported and it is in the past. + */ + [[nodiscard]] MCP_API bool secret_expired(std::int64_t now_seconds) const; +}; + +/** + * @brief Deserialize client information from an RFC 7591 registration response. + * + * @param json JSON registration response. + * @param information Structure to populate. + * + * @note `issuer` and `source` are not carried by the wire format; the caller sets them from the + * metadata document the registration endpoint was taken from. + */ +MCP_API void from_json(const nlohmann::json& json, OAuthClientInformation& information); + +/** + * @brief Serialize client information, including its issuer binding. + * + * @param json JSON object to populate. + * @param information Information to serialize. + */ +MCP_API void to_json(nlohmann::json& json, const OAuthClientInformation& information); + +/** + * @brief Storage for acquired client credentials, keyed by authorization-server issuer. + * + * @details Separate from TokenStore, which is keyed by MCP server URL: credentials are bound to one + * authorization server (SEP-2352), so they are keyed by issuer. + */ +class ClientCredentialStore { + public: + virtual ~ClientCredentialStore() = default; + + /** + * @brief Store or replace the credentials associated with an issuer. + * + * @param issuer Authorization server issuer used as the storage key. + * @param information Credentials to persist. + */ + virtual void store(const std::string& issuer, OAuthClientInformation information) = 0; + + /** + * @brief Load the credentials associated with an issuer. + * + * @param issuer Authorization server issuer used as the storage key. + * @return The stored credentials, if present. + */ + [[nodiscard]] virtual std::optional load( + const std::string& issuer) const = 0; + + /** + * @brief Remove the credentials associated with an issuer. + * + * @param issuer Authorization server issuer used as the storage key. + */ + virtual void remove(const std::string& issuer) = 0; +}; + +/** + * @brief Thread-safe in-memory client credential store. + */ +class MCP_API InMemoryClientCredentialStore : public ClientCredentialStore { + public: + InMemoryClientCredentialStore(); + ~InMemoryClientCredentialStore() override; + + InMemoryClientCredentialStore(const InMemoryClientCredentialStore&) = delete; + InMemoryClientCredentialStore& operator=(const InMemoryClientCredentialStore&) = delete; + InMemoryClientCredentialStore(InMemoryClientCredentialStore&&) = delete; + InMemoryClientCredentialStore& operator=(InMemoryClientCredentialStore&&) = delete; + + /** + * @brief Store or replace credentials in the in-memory cache. + * + * @param issuer Authorization server issuer used as the storage key. + * @param information Credentials to persist. + */ + void store(const std::string& issuer, OAuthClientInformation information) override; + + /** + * @brief Load credentials from the in-memory cache. + * + * @param issuer Authorization server issuer used as the storage key. + * @return The stored credentials, if present. + */ + std::optional load(const std::string& issuer) const override; + + /** + * @brief Remove credentials from the in-memory cache. + * + * @param issuer Authorization server issuer used as the storage key. + */ + void remove(const std::string& issuer) override; + + private: + struct Impl; + std::unique_ptr impl_; +}; + +/** + * @brief Application-supplied inputs to client identity selection. + */ +struct ClientIdentityConfig { + /// Credentials the application obtained out of band. When present, registration is never + /// attempted. Set `issuer` on them to say which authorization server they belong to: it is + /// required whenever `client_secret` is set, because an unbound secret is refused rather than + /// presented to whichever server the protected-resource document named. See + /// `select_client_identity` for the full rule. + std::optional pre_registered; + /// URL of a client ID metadata document the application publishes. Used as the client + /// identifier itself when the authorization server advertises support for it. + std::optional client_metadata_url; + /// Metadata sent when dynamic registration is the remaining path. + OAuthClientMetadata metadata; +}; + +/** + * @brief The authorization-server facts client identity selection depends on. + */ +struct ClientIdentityServerFacts { + std::string issuer; ///< Issuer recorded from the metadata. + bool client_id_metadata_document_supported{false}; ///< `client_id_metadata_document_supported`. + std::optional registration_endpoint; ///< RFC 7591 registration endpoint, if any. + std::vector scopes_supported; ///< `scopes_supported`, used for offline access. +}; + +/** + * @brief The client identity path chosen for one authorization server. + */ +enum class ClientIdentityDecision { + use_pre_registered, ///< Present the application's injected credentials. + use_client_id_metadata_document, ///< Present the configured document URL as the client ID. + reuse_stored_registration, ///< Present credentials already registered with this issuer. + register_dynamically, ///< Register with this issuer and present the result. + unavailable, ///< No path is available; authorization cannot proceed. +}; + +/** + * @brief Short, stable description of a client identity decision. + * + * @param decision Decision to describe. + * @return A description suitable for diagnostics. + */ +[[nodiscard]] MCP_API std::string_view describe(ClientIdentityDecision decision); + +/** + * @brief Choose the client identity path for one authorization server, without any network access. + * + * @param config Application-supplied identity inputs. + * @param server Facts recorded from the selected authorization server's metadata. + * @param stored Credentials already held for this issuer, if any. + * @return The chosen path. + * + * @details Precedence, highest first: + * + * 1. Injected credentials, with no fallback. Their issuer binding decides: + * - `pre_registered.issuer` names the issuer being contacted: the credentials are used. + * - It names a different issuer: `unavailable`. + * - It is empty: a public client (no `client_secret`) is used; a `client_secret` is refused with + * `unavailable` until `pre_registered.issuer` (or `OAuthAuthorizationConfig::client_issuer` on + * the shorthand path) is set. + * 2. A configured client ID metadata document URL, when the server advertises support for one. + * Dynamic registration is skipped on this path. + * 3. Credentials stored for this exact issuer; an entry recorded under another issuer is ignored. + * 4. Dynamic registration (RFC 7591, deprecated in favour of client ID metadata documents), when + * the server publishes a registration endpoint. + */ +[[nodiscard]] MCP_API ClientIdentityDecision +select_client_identity(const ClientIdentityConfig& config, const ClientIdentityServerFacts& server, + const std::optional& stored); + +/** + * @brief Build the RFC 7591 dynamic client registration request body. + * + * @param metadata Client metadata to register. + * @param server Facts recorded from the authorization server's metadata. + * @return The JSON request body. + * + * @details Carries `application_type` (SEP-837) unconditionally. `offline_access` is appended to the + * requested scope only when the server advertises it in `scopes_supported`. + */ +[[nodiscard]] MCP_API nlohmann::json build_registration_request( + const OAuthClientMetadata& metadata, const ClientIdentityServerFacts& server); + +} // namespace mcp::auth diff --git a/include/mcp/auth/metadata_policy.hpp b/include/mcp/auth/metadata_policy.hpp new file mode 100644 index 0000000..c1538f2 --- /dev/null +++ b/include/mcp/auth/metadata_policy.hpp @@ -0,0 +1,201 @@ +#pragma once + +#include + +#include +#include +#include +#include +#include +#include + +namespace mcp::auth { + +namespace constants { + +constexpr std::size_t g_default_max_metadata_bytes = 256 * 1024; +constexpr std::size_t g_default_max_metadata_redirects = 3; + +} // namespace constants + +namespace detail { + +/** + * @brief Bound and clean peer-controlled text before embedding it in a diagnostic message. + * + * @param value Text that came from a peer: a URL, a host, an issuer, a metadata field, a response + * body. Anything the SDK did not author itself. + * @return Well-formed UTF-8, bounded by a fixed budget and ending in an ellipsis when it was cut, + * with every character that could end a line replaced by a space and every ill-formed byte + * replaced by `?`. + * + * @details Keeps a peer-chosen value from forging a log line with embedded newlines or flooding the + * message. Truncation stops on a codepoint boundary, so a multi-byte character straddling the budget + * is dropped whole. + */ +[[nodiscard]] MCP_API std::string sanitize_for_diagnostics(std::string_view value); + +} // namespace detail + +/** + * @brief Outcome of validating a metadata target before any network access occurs. + * + * @details Every value other than `allowed` means the SDK refused the target. Refusal happens + * before host resolution, so a refused target is never contacted. + */ +enum class MetadataUrlDecision { + allowed, ///< Target passed every configured control. + malformed_url, ///< URL could not be decomposed into scheme, host and port, including a + ///< port that is not a plain decimal number or a host that is nothing but dots. + scheme_not_allowed, ///< Scheme is not `https`, and the plain-HTTP loopback opt-out did not apply. + origin_denied, ///< Origin matched the application's deny list. + origin_not_allowed, ///< Origin is not present in the application's allow list. + address_link_local, ///< Address is IPv4 link-local (169.254.0.0/16) or IPv6 link-local. + address_private, ///< Address is in an RFC 1918 or IPv6 unique-local range. + address_loopback, ///< Address is loopback and the loopback opt-out is not enabled. + address_multicast, ///< Address is multicast or broadcast. + address_reserved, ///< Address is otherwise reserved, unspecified or non-routable. + redirect_limit_exceeded, ///< Redirect chain exceeded the configured bound. + response_too_large, ///< Response body exceeded the configured size cap. + denied_origin_entry_malformed, ///< A `denied_origins` entry is not a bare origin, so the rule + ///< it states cannot be applied. Reported against the offending + ///< entry, not against the target. +}; + +/** + * @brief Human-readable description of a metadata URL decision. + * + * @param decision Decision to describe. + * @return A short, stable description suitable for diagnostics. + */ +[[nodiscard]] MCP_API std::string_view describe(MetadataUrlDecision decision); + +/** + * @brief Application-controlled policy governing outbound OAuth metadata requests. + * + * @details Consulted before any OAuth metadata, protected-resource or token request is issued; a + * challenge-supplied `resource_metadata` URL is attacker-influenced input. + * + * A default-constructed policy denies every origin: `allowed_origins` is empty, and an origin that is + * not listed is refused. + */ +struct MetadataFetchPolicy { + /// Origins the application permits, each written as `scheme://host[:port]` with nothing after the + /// authority. An empty list denies every origin. Compared after canonicalizing scheme and host + /// case, trailing host dots, an explicit port equal to the scheme's default (`:443` for `https`, + /// `:80` for `http`, also when written as `:0443`), and an IP literal's textual form (expanded + /// and compressed IPv6 spellings compare equal). An entry that carries a path, query or fragment, + /// or whose port is not a plain in-range decimal number, matches nothing. + std::vector allowed_origins; + + /// Origins the application refuses. Consulted before `allowed_origins`, so a denied origin is + /// refused even when it also appears in the allow list. Compared with the same canonicalization + /// as `allowed_origins`. + /// + /// Every entry must be a bare origin. Unlike `allowed_origins`, an entry that is not one -- it + /// carries a path (a lone trailing `/` included), a query or a fragment, or its port is not a + /// plain in-range decimal number -- is rejected rather than ignored: `validate_metadata_url` + /// throws `MetadataPolicyError` with `MetadataUrlDecision::denied_origin_entry_malformed`, naming + /// the offending entry, and refuses every target until the entry is corrected. + std::vector denied_origins; + + /// Consulted only for an origin `allowed_origins` does not list; returning true admits it. The + /// origin passed to the callback is the canonicalized form (see `allowed_origins`), not the raw + /// text of the URL. + /// + /// The hook can only widen: `denied_origins` is consulted first and still refuses, and the + /// decision is made before host resolution, so an origin this hook rejects is never contacted. + std::function origin_allowance; + + /// Permit plain `http://` and loopback addresses. This is a narrow opt-out for loopback + /// development and test fixtures only; it is never enabled implicitly and it does not relax + /// any control other than the HTTPS requirement and the loopback address block. + bool allow_plain_http_loopback{false}; + + /// Maximum number of response body bytes accepted from a metadata or token endpoint. A + /// response that exceeds the cap is rejected without being fully buffered. + std::size_t max_response_bytes{constants::g_default_max_metadata_bytes}; + + /// Maximum number of HTTP redirects followed. Each redirect target is validated afresh. + std::size_t max_redirects{constants::g_default_max_metadata_redirects}; +}; + +/** + * @brief Extract the origin of a URL as `scheme://host[:port]`. + * + * @param url Absolute URL to decompose. + * @return The origin exactly as written in the URL, or an empty string when the URL is malformed. + * + * @note No normalization is performed: case, an explicit default port and a trailing dot are all + * preserved, so allow-list entries must be written the way the URL will appear. + */ +[[nodiscard]] MCP_API std::string metadata_url_origin(const std::string& url); + +/** + * @brief Validate a metadata target URL against the scheme and origin controls. + * + * @param policy Application policy to apply. + * @param url Absolute URL that the SDK is about to fetch. + * @return `MetadataUrlDecision::allowed` when the URL may be resolved, otherwise the refusal + * reason. + * + * @throws MetadataPolicyError With `MetadataUrlDecision::denied_origin_entry_malformed` when any + * `MetadataFetchPolicy::denied_origins` entry is not a bare origin. The deny list is + * examined first, so such a policy refuses every target. The error names the offending + * entry rather than the URL. + * + * @details Runs before host resolution: when it refuses, the host is never resolved and no socket is + * opened. A host written as an IP literal is classified here, so a URL naming `169.254.169.254` or an + * RFC 1918 address is refused without any lookup. The origin derived from `url` is canonicalized (see + * `MetadataFetchPolicy::allowed_origins`) before the deny list, allow list or `origin_allowance` sees + * it, and the https-required check and loopback opt-out run on the same canonical scheme and host; a + * URL whose port or host does not canonicalize is refused as `malformed_url`. `metadata_url_origin` + * is unaffected and never normalizes its result. + */ +[[nodiscard]] MCP_API MetadataUrlDecision validate_metadata_url(const MetadataFetchPolicy& policy, + const std::string& url); + +/** + * @brief Classify a resolved address against the SSRF controls. + * + * @param policy Application policy to apply. + * @param address_literal Resolved address in textual form. + * @return `MetadataUrlDecision::allowed` when the address may be connected to, otherwise the + * refusal reason. + * + * @details Applied to every address produced by resolution. The SDK connects only to addresses that + * pass, and pins the addresses from that single resolution rather than resolving again. + */ +[[nodiscard]] MCP_API MetadataUrlDecision validate_metadata_address(const MetadataFetchPolicy& policy, + const std::string& address_literal); + +/** + * @brief Error raised when a metadata target is refused by the fetch policy. + * + * @details Thrown instead of performing the request: for a URL or origin decision the host is never + * resolved, and for an address decision the address is never connected to. For + * `MetadataUrlDecision::denied_origin_entry_malformed`, `target()` is the offending `denied_origins` + * entry rather than the request target. + */ +class MCP_API MetadataPolicyError : public std::runtime_error { + public: + /** + * @brief Construct a policy refusal. + * + * @param decision Reason the target was refused. + * @param target URL or address that was refused. + */ + MetadataPolicyError(MetadataUrlDecision decision, std::string target); + + /** @return Reason the target was refused. */ + [[nodiscard]] MetadataUrlDecision decision() const noexcept { return decision_; } + + /** @return URL or address that was refused. */ + [[nodiscard]] const std::string& target() const noexcept { return target_; } + + private: + MetadataUrlDecision decision_; + std::string target_; +}; + +} // namespace mcp::auth diff --git a/include/mcp/auth/oauth.hpp b/include/mcp/auth/oauth.hpp index b23a8b0..e9ba33e 100644 --- a/include/mcp/auth/oauth.hpp +++ b/include/mcp/auth/oauth.hpp @@ -1,40 +1,31 @@ #pragma once +#include +#include +#include #include #include #include +#include #include #include #include -#include -#include -#include #include #include -#include -#include -#include -#include -#include -#include #include +#include +#include #include #include -#include #include #include -#include #include #include -#include -#include #include namespace mcp::auth { -namespace beast = boost::beast; -namespace http = beast::http; namespace net = boost::asio; namespace constants { @@ -57,109 +48,15 @@ constexpr std::size_t g_default_cache_ttl_seconds = 300; } // namespace constants namespace detail { -inline std::string base64_encode(const unsigned char* data, std::size_t len) { - std::string result; - result.reserve(((len + 2) / 3) * 4); - - for (std::size_t i = 0; i < len; i += 3) { - unsigned int n = static_cast(data[i]) << constants::g_shift16; - if (i + 1 < len) { - n |= static_cast(data[i + 1]) << constants::g_shift8; - } - if (i + 2 < len) { - n |= static_cast(data[i + 2]); - } - - result.push_back( - mcp::constants::g_alphabet[(n >> constants::g_shift18) & constants::g_mask0x3F]); - result.push_back( - mcp::constants::g_alphabet[(n >> constants::g_shift12) & constants::g_mask0x3F]); - result.push_back( - (i + 1 < len) - ? mcp::constants::g_alphabet[(n >> constants::g_shift6) & constants::g_mask0x3F] - : '='); - result.push_back((i + 2 < len) ? mcp::constants::g_alphabet[n & constants::g_mask0x3F] : '='); - } - - return result; -} - -inline std::string base64url_encode(const unsigned char* data, std::size_t len) { - auto encoded = base64_encode(data, len); - - for (auto& ch : encoded) { - if (ch == '+') { - ch = '-'; - } else if (ch == '/') { - ch = '_'; - } - } - encoded.erase(std::remove(encoded.begin(), encoded.end(), '='), encoded.end()); - return encoded; -} +MCP_API std::string base64_encode(const unsigned char* data, std::size_t len); +MCP_API std::string base64url_encode(const unsigned char* data, std::size_t len); /// SHA-256 via OpenSSL EVP interface. -inline std::array sha256(const std::string& input) { - std::array digest{}; - unsigned int digest_len = 0; - - std::unique_ptr ctx(EVP_MD_CTX_new(), EVP_MD_CTX_free); - if (!ctx || EVP_DigestInit_ex(ctx.get(), EVP_sha256(), nullptr) != 1 || - EVP_DigestUpdate(ctx.get(), input.data(), input.size()) != 1 || - EVP_DigestFinal_ex(ctx.get(), digest.data(), &digest_len) != 1) { - throw std::runtime_error("OpenSSL SHA-256 failed"); - } +MCP_API std::array sha256(const std::string& input); - return digest; -} - -inline std::string generate_random_string(std::size_t length) { - static const std::size_t s_charset_size = mcp::constants::g_unreserved_chars.size(); - // Largest multiple of charset_size that fits in a byte (avoids modulo bias) - static const auto s_bias_limit = - static_cast((256 / s_charset_size) * s_charset_size); - - std::string result; - result.reserve(length); - while (result.size() < length) { - unsigned char byte = 0; - if (RAND_bytes(&byte, 1) != 1) { - throw std::runtime_error("RAND_bytes failed"); - } - if (byte < s_bias_limit) { - result.push_back(mcp::constants::g_unreserved_chars[byte % s_charset_size]); - } - } - return result; -} - -inline std::string url_encode(const std::string& value) { - std::string result; - result.reserve(value.size() * 3); - - for (unsigned char ch : value) { - if ((ch >= 'A' && ch <= 'Z') || (ch >= 'a' && ch <= 'z') || (ch >= '0' && ch <= '9') || - ch == '-' || ch == '_' || ch == '.' || ch == '~') { - result.push_back(static_cast(ch)); - } else { - result.push_back('%'); - result.push_back(mcp::constants::g_hex_digits_upper[ch >> constants::g_shift4]); - result.push_back(mcp::constants::g_hex_digits_upper[ch & constants::g_mask0x0F]); - } - } - return result; -} - -inline std::string build_form_body(const KeyValuePairList& params) { - std::string body; - for (const auto& [key, value] : params) { - if (!body.empty()) { - body.push_back('&'); - } - body += url_encode(key) + "=" + url_encode(value); - } - return body; -} +MCP_API std::string generate_random_string(std::size_t length); +MCP_API std::string url_encode(const std::string& value); +MCP_API std::string build_form_body(const KeyValuePairList& params); } // namespace detail @@ -178,21 +75,7 @@ struct PkcePair { * @param verifier_length Length of the verifier string, between 43 and 128 characters. * @return A PKCE pair suitable for OAuth 2.1 authorization code flows. */ -inline PkcePair generate_pkce_pair(std::size_t verifier_length = constants::g_default_verifier_length) { - if (verifier_length < constants::g_min_verifier_length || - verifier_length > constants::g_max_verifier_length) { - throw std::invalid_argument("PKCE verifier length must be 43-128 characters"); - } - - PkcePair pair; - pair.code_verifier = detail::generate_random_string(verifier_length); - pair.challenge_method = "S256"; - - auto hash = detail::sha256(pair.code_verifier); - pair.code_challenge = detail::base64url_encode(hash.data(), hash.size()); - - return pair; -} +MCP_API PkcePair generate_pkce_pair(std::size_t verifier_length = constants::g_default_verifier_length); /** * @brief OAuth token response data returned by an authorization server. @@ -212,14 +95,8 @@ struct TokenResponse { * @param margin Safety margin in seconds applied before reported expiry. * @return True when the token is expired or within the safety margin. */ - [[nodiscard]] bool is_expired( - int margin = constants::g_default_token_lifetime_safety_margin_seconds) const { - if (!expires_in.has_value()) { - return false; - } - auto expiry = received_at + std::chrono::seconds(*expires_in) - std::chrono::seconds(margin); - return std::chrono::steady_clock::now() >= expiry; - } + [[nodiscard]] MCP_API bool is_expired( + int margin = constants::g_default_token_lifetime_safety_margin_seconds) const; }; /** @@ -228,20 +105,7 @@ struct TokenResponse { * @param j JSON token payload. * @param t Token response to populate. */ -inline void from_json(const nlohmann::json& j, TokenResponse& t) { - j.at("access_token").get_to(t.access_token); - t.token_type = j.value("token_type", "Bearer"); - if (j.contains("refresh_token")) { - t.refresh_token = j.at("refresh_token").get(); - } - if (j.contains("expires_in")) { - t.expires_in = j.at("expires_in").get(); - } - if (j.contains("scope")) { - t.scope = j.at("scope").get(); - } - t.received_at = std::chrono::steady_clock::now(); -} +MCP_API void from_json(const nlohmann::json& j, TokenResponse& t); /** * @brief Serialize a token response to JSON. @@ -249,18 +113,7 @@ inline void from_json(const nlohmann::json& j, TokenResponse& t) { * @param j JSON object to populate. * @param t Token response to serialize. */ -inline void to_json(nlohmann::json& j, const TokenResponse& t) { - j = nlohmann::json{{"access_token", t.access_token}, {"token_type", t.token_type}}; - if (t.refresh_token) { - j["refresh_token"] = *t.refresh_token; - } - if (t.expires_in) { - j["expires_in"] = *t.expires_in; - } - if (t.scope) { - j["scope"] = *t.scope; - } -} +MCP_API void to_json(nlohmann::json& j, const TokenResponse& t); /** * @brief Abstract storage interface for OAuth tokens keyed by server URL. @@ -293,18 +146,23 @@ class TokenStore { /** * @brief Thread-safe in-memory token store implementation. */ -class InMemoryTokenStore : public TokenStore { +class MCP_API InMemoryTokenStore : public TokenStore { public: + InMemoryTokenStore(); + ~InMemoryTokenStore() override; + + InMemoryTokenStore(const InMemoryTokenStore&) = delete; + InMemoryTokenStore& operator=(const InMemoryTokenStore&) = delete; + InMemoryTokenStore(InMemoryTokenStore&&) = delete; + InMemoryTokenStore& operator=(InMemoryTokenStore&&) = delete; + /** * @brief Store or replace a token in the in-memory cache. * * @param server_url MCP server URL used as the storage key. * @param token Token data to persist. */ - void store(const std::string& server_url, TokenResponse token) override { - std::lock_guard lock(mutex_); - tokens_[server_url] = std::move(token); - } + void store(const std::string& server_url, TokenResponse token) override; /** * @brief Load a token from the in-memory cache. @@ -312,28 +170,18 @@ class InMemoryTokenStore : public TokenStore { * @param server_url MCP server URL used as the storage key. * @return The stored token, if present. */ - std::optional load(const std::string& server_url) const override { - std::lock_guard lock(mutex_); - auto it = tokens_.find(server_url); - if (it == tokens_.end()) { - return std::nullopt; - } - return it->second; - } + std::optional load(const std::string& server_url) const override; /** * @brief Remove a token from the in-memory cache. * * @param server_url MCP server URL used as the storage key. */ - void remove(const std::string& server_url) override { - std::lock_guard lock(mutex_); - tokens_.erase(server_url); - } + void remove(const std::string& server_url) override; private: - mutable std::mutex mutex_; - std::unordered_map tokens_; + struct Impl; + std::unique_ptr impl_; }; /** @@ -348,20 +196,96 @@ struct OAuthConfig { std::string redirect_uri; ///< Redirect URI used during authorization code flow. std::optional scope; ///< Optional requested scope string. std::optional resource; ///< Optional resource or audience hint. + /// Client authentication method for the token endpoint, spelled as the authorization server + /// spells it in `token_endpoint_auth_methods_supported`. `client_secret_basic` puts the + /// credentials in the HTTP Basic header and nowhere else; `none` sends no secret at all; + /// anything else, including leaving this unset, puts the secret in the request body. + std::optional token_endpoint_auth_method; }; +/** + * @brief Resolves a host and port to candidate address literals. + * + * @details When the fetch policy refuses a target the resolver is never invoked. When none is + * installed the executor's system resolver is used. + */ +using HostResolver = + std::function(const std::string& host, const std::string& port)>; + /** * @brief Minimal HTTP client for OAuth token exchange and metadata retrieval. */ -class OAuthHttpClient { +class OAuthHttpClientScope; + +namespace detail { + +/// One scope's abort latch. Opaque here and defined in the implementation: a scope holds its own +/// latch rather than a name for one the client keeps, which is what lets the latch die with the +/// scope instead of accumulating on the client for the life of the process. +struct OAuthScopeState; + +/// Reach-in for this SDK's own tests, declared opaque on purpose. +/// +/// It hands out nothing callable: the type is defined only inside the implementation, and the +/// accessors that use it are declared in a header that is not installed. A friend declaration does +/// not affect layout, so binary compatibility is unaffected. +struct OAuthTestAccess; + +} // namespace detail + +class MCP_API OAuthHttpClient { public: /** * @brief Construct an OAuth HTTP client. * * @param executor Executor used for asynchronous operations. + * + * @note Without an explicit fetch policy the client holds a default-constructed + * `MetadataFetchPolicy` (empty origin allow list, https-only, non-loopback), so every request is + * refused until a policy is installed via set_metadata_policy(). + */ + explicit OAuthHttpClient(const net::any_io_executor& executor); + + /** + * @brief Apply an outbound-request policy to every request this client issues. + * + * @param policy Policy governing schemes, origins, resolved addresses, response size and + * redirect depth. + * + * @details May be called at any time: an exchange already running keeps the policy it started + * with, and the next one picks up the new value. Each request is validated before host + * resolution, every resolved address is classified before connecting and pinned for the + * connection, response bodies are capped, and redirects are bounded and individually + * re-validated. + * + * @warning Do not call this and set_host_resolver() in sequence to change both: an exchange + * started between the two calls runs with one new value and one old one. Use configure(). + */ + void set_metadata_policy(MetadataFetchPolicy policy); + + /** + * @brief Install a custom host resolver. + * + * @param resolver Resolver invoked in place of the system resolver. + * + * @warning Do not call this and set_metadata_policy() in sequence to change both; use + * configure(). + */ + void set_host_resolver(HostResolver resolver); + + /** + * @brief Install an outbound-request policy and a host resolver as one indivisible change. + * + * @param policy Policy governing schemes, origins, resolved addresses, response size and + * redirect depth. + * @param resolver Resolver invoked in place of the system resolver; an empty resolver restores + * the executor's system resolver. + * + * @details An exchange started around this call runs entirely with the configuration before it or + * entirely with the configuration after it. Use this whenever both are changed on a client that + * may already be serving requests. */ - explicit OAuthHttpClient(const net::any_io_executor& executor) - : strand_(net::make_strand(executor)) {} + void configure(MetadataFetchPolicy policy, HostResolver resolver); /** * @brief Exchange an authorization code for an access token. @@ -372,22 +296,7 @@ class OAuthHttpClient { * @return A task resolving to the parsed token response. */ Task exchange_code(const OAuthConfig& config, const std::string& code, - const std::string& code_verifier) { - KeyValuePairList params = { - {"grant_type", "authorization_code"}, {"code", code}, - {"redirect_uri", config.redirect_uri}, {"client_id", config.client_id}, - {"code_verifier", code_verifier}, - }; - - if (config.client_secret) { - params.emplace_back("client_secret", *config.client_secret); - } - if (config.resource) { - params.emplace_back("resource", *config.resource); - } - - co_return co_await post_token_request(config.token_endpoint, params); - } + const std::string& code_verifier); /** * @brief Refresh an access token using a refresh token. @@ -396,22 +305,7 @@ class OAuthHttpClient { * @param refresh_token Refresh token issued by the authorization server. * @return A task resolving to the parsed token response. */ - Task refresh_token(const OAuthConfig& config, const std::string& refresh_token) { - KeyValuePairList params = { - {"grant_type", "refresh_token"}, - {"refresh_token", refresh_token}, - {"client_id", config.client_id}, - }; - - if (config.client_secret) { - params.emplace_back("client_secret", *config.client_secret); - } - if (config.resource) { - params.emplace_back("resource", *config.resource); - } - - co_return co_await post_token_request(config.token_endpoint, params); - } + Task refresh_token(const OAuthConfig& config, const std::string& refresh_token); /** * @brief Fetch a JSON document from an OAuth discovery endpoint. @@ -419,122 +313,96 @@ class OAuthHttpClient { * @param url HTTP URL to fetch. * @return A task resolving to the parsed JSON body. */ - Task get_json(const std::string& url) { - auto parsed = parse_url(url); - - co_await net::post(strand_, net::use_awaitable); - - net::ip::tcp::resolver resolver(strand_); - auto endpoints = co_await resolver.async_resolve(parsed.host, parsed.port, net::use_awaitable); - - beast::tcp_stream stream(strand_); - stream.expires_after(std::chrono::seconds(mcp::constants::g_http_timeout_seconds)); - co_await stream.async_connect(endpoints, net::use_awaitable); - - http::request req{http::verb::get, parsed.path, - mcp::constants::g_http_version_11}; - req.set(http::field::host, parsed.host); - req.set(http::field::accept, "application/json"); + Task get_json(const std::string& url); - stream.expires_after(std::chrono::seconds(mcp::constants::g_http_timeout_seconds)); - co_await http::async_write(stream, req, net::use_awaitable); - - beast::flat_buffer buffer; - http::response res; - co_await http::async_read(stream, buffer, res, net::use_awaitable); - - beast::error_code ec; - (void)stream.socket().shutdown(net::ip::tcp::socket::shutdown_both, ec); - - if (res.result_int() >= mcp::constants::g_http_bad_request) { - throw std::runtime_error("HTTP GET " + url + " failed with status " + - std::to_string(res.result_int())); - } + /** + * @brief Post a JSON document to an OAuth endpoint and read the JSON reply. + * + * @param url HTTP URL to post to. + * @param body JSON request body. + * @return A task resolving to the parsed JSON response. + * + * @details Used for RFC 7591 dynamic client registration. The target is validated against the + * fetch policy exactly like every other request this client issues, and redirects are not + * followed for a POST. + */ + Task post_json(const std::string& url, const nlohmann::json& body); - auto json = nlohmann::json::parse(res.body(), nullptr, false); - if (json.is_discarded()) { - throw std::runtime_error("Failed to parse JSON from " + url); - } + /** + * @brief Abort every HTTP exchange currently in flight on this client, and refuse every one + * issued afterward. + * + * @details Closes the underlying socket of each active exchange, so a pending resolve, connect, + * write, or read completes with an error instead of hanging. Safe to call from any thread. Sticky + * and irreversible: a request issued afterward -- even one that has not yet made its first + * network call -- fails immediately instead of running to completion. Idempotent. + */ + void abort_pending(); - co_return json; - } + /** + * @brief Open an independently abortable scope on this client. + * + * @return A scope that issues requests through this client but can be aborted on its own. + * + * @details Requests issued through a scope are tracked against it; aborting it closes those and + * refuses later ones, leaving every other scope and the client's own unscoped requests untouched. + * `abort_pending()` still ends everything, scopes included. A scope keeps the underlying client + * alive, so it stays usable even if the `OAuthHttpClient` object it came from is destroyed. + */ + [[nodiscard]] OAuthHttpClientScope make_scope(); private: - struct ParsedUrl { - std::string host; - std::string port; - std::string path; - }; - - static ParsedUrl parse_url(const std::string& url) { - if (!url.starts_with(mcp::constants::g_http_prefix)) { - throw std::invalid_argument("OAuth HTTP client URL must start with http://"); - } - - auto authority_and_path = url.substr(mcp::constants::g_http_prefix.size()); - auto path_sep = authority_and_path.find('/'); - auto authority = authority_and_path.substr(0, path_sep); - auto path_val = path_sep == std::string::npos ? "/" : authority_and_path.substr(path_sep); - - std::string host; - std::string port = "80"; - auto colon = authority.find(':'); - if (colon == std::string::npos) { - host = std::move(authority); - } else { - host = authority.substr(0, colon); - port = authority.substr(colon + 1); - } - - return {std::move(host), std::move(port), std::move(path_val)}; - } - - Task post_token_request(const std::string& token_endpoint, - const KeyValuePairList& params) { - auto parsed = parse_url(token_endpoint); - auto form_body = detail::build_form_body(params); - - co_await net::post(strand_, net::use_awaitable); - - net::ip::tcp::resolver resolver(strand_); - auto endpoints = co_await resolver.async_resolve(parsed.host, parsed.port, net::use_awaitable); + struct Impl; + std::shared_ptr impl_; - beast::tcp_stream stream(strand_); - stream.expires_after(std::chrono::seconds(mcp::constants::g_http_timeout_seconds)); - co_await stream.async_connect(endpoints, net::use_awaitable); + friend class OAuthHttpClientScope; + friend struct detail::OAuthTestAccess; +}; - http::request req{http::verb::post, parsed.path, - mcp::constants::g_http_version_11}; - req.set(http::field::host, parsed.host); - req.set(http::field::content_type, "application/x-www-form-urlencoded"); - req.set(http::field::accept, "application/json"); - req.body() = std::move(form_body); - req.prepare_payload(); +/** + * @brief An independently abortable view of an OAuthHttpClient. + * + * @details Issues requests exactly as the client does, but tracks them separately so `abort()` ends + * this scope's work alone. Obtained from `OAuthHttpClient::make_scope()`. Use one wherever the client + * is shared: `OAuthHttpClient::abort_pending()` disables the client permanently for every holder. + * + * Copyable, and every copy names the same scope, so aborting through any copy aborts them all. + */ +class MCP_API OAuthHttpClientScope { + public: + /** @see OAuthHttpClient::exchange_code */ + Task exchange_code(const OAuthConfig& config, const std::string& code, + const std::string& code_verifier); - stream.expires_after(std::chrono::seconds(mcp::constants::g_http_timeout_seconds)); - co_await http::async_write(stream, req, net::use_awaitable); + /** @see OAuthHttpClient::refresh_token */ + Task refresh_token(const OAuthConfig& config, const std::string& refresh_token); - beast::flat_buffer buffer; - http::response res; - co_await http::async_read(stream, buffer, res, net::use_awaitable); + /** @see OAuthHttpClient::get_json */ + Task get_json(const std::string& url); - beast::error_code ec; - (void)stream.socket().shutdown(net::ip::tcp::socket::shutdown_both, ec); + /** @see OAuthHttpClient::post_json */ + Task post_json(const std::string& url, const nlohmann::json& body); - if (res.result_int() >= mcp::constants::g_http_bad_request) { - throw std::runtime_error("Token request failed with status " + - std::to_string(res.result_int()) + ": " + res.body()); - } + /** + * @brief Abort this scope's in-flight exchanges and refuse every one issued through it after. + * + * @details The scoped counterpart of `OAuthHttpClient::abort_pending()`: safe from any thread, + * sticky, irreversible and idempotent. Requests issued through the client directly, or through + * any other scope, are unaffected. + */ + void abort(); - auto response_json = nlohmann::json::parse(res.body(), nullptr, false); - if (response_json.is_discarded()) { - throw std::runtime_error("Failed to parse token response JSON"); - } + private: + friend class OAuthHttpClient; - co_return response_json.get(); - } + OAuthHttpClientScope(std::shared_ptr impl, + std::shared_ptr state); - net::strand strand_; + std::shared_ptr impl_; + /// The latch itself, not a name for one held elsewhere. Every exchange issued through this + /// scope holds the same control block, so the flag lives exactly as long as someone can + /// still consult it and is freed once nobody can. + std::shared_ptr state_; }; /** @@ -553,18 +421,7 @@ struct ProtectedResourceMetadata { * @param j JSON metadata payload. * @param m Metadata structure to populate. */ -inline void from_json(const nlohmann::json& j, ProtectedResourceMetadata& m) { - m.raw = j; - if (j.contains("resource")) { - j.at("resource").get_to(m.resource); - } - if (j.contains("authorization_servers")) { - j.at("authorization_servers").get_to(m.authorization_servers); - } - if (j.contains("scopes_supported")) { - m.scopes_supported = j.at("scopes_supported").get>(); - } -} +MCP_API void from_json(const nlohmann::json& j, ProtectedResourceMetadata& m); /** * @brief Metadata exposed by an OAuth authorization server. @@ -581,7 +438,13 @@ struct AuthServerMetadata { std::optional> grant_types_supported; ///< Optional supported grant types. std::optional> code_challenge_methods_supported; ///< Optional PKCE methods. - nlohmann::json raw; ///< Raw source document. + std::optional + authorization_response_iss_parameter_supported; ///< Optional RFC 9207 `iss` support flag. + std::optional + client_id_metadata_document_supported; ///< Optional client ID metadata document support flag. + std::optional> + token_endpoint_auth_methods_supported; ///< Optional token endpoint auth methods. + nlohmann::json raw; ///< Raw source document. }; /** @@ -590,37 +453,7 @@ struct AuthServerMetadata { * @param j JSON metadata payload. * @param m Metadata structure to populate. */ -inline void from_json(const nlohmann::json& j, AuthServerMetadata& m) { - m.raw = j; - if (j.contains("issuer")) { - j.at("issuer").get_to(m.issuer); - } - if (j.contains("authorization_endpoint")) { - j.at("authorization_endpoint").get_to(m.authorization_endpoint); - } - if (j.contains("token_endpoint")) { - j.at("token_endpoint").get_to(m.token_endpoint); - } - if (j.contains("revocation_endpoint")) { - m.revocation_endpoint = j.at("revocation_endpoint").get(); - } - if (j.contains("registration_endpoint")) { - m.registration_endpoint = j.at("registration_endpoint").get(); - } - if (j.contains("scopes_supported")) { - m.scopes_supported = j.at("scopes_supported").get>(); - } - if (j.contains("response_types_supported")) { - m.response_types_supported = j.at("response_types_supported").get>(); - } - if (j.contains("grant_types_supported")) { - m.grant_types_supported = j.at("grant_types_supported").get>(); - } - if (j.contains("code_challenge_methods_supported")) { - m.code_challenge_methods_supported = - j.at("code_challenge_methods_supported").get>(); - } -} +MCP_API void from_json(const nlohmann::json& j, AuthServerMetadata& m); template struct CachedEntry { @@ -630,10 +463,19 @@ struct CachedEntry { [[nodiscard]] bool is_expired() const { return std::chrono::steady_clock::now() >= expires_at; } }; +/** + * @brief Caller's verdict on a protected-resource document, before it is trusted or cached. + * + * @details Invoked with the parsed document on the way out of discovery, whether it was just fetched + * or served from the cache. Throwing rejects it: the exception reaches the caller, no cache entry is + * written, and discovery does not fall through to the next candidate URL. + */ +using ProtectedResourceAcceptor = std::function; + /** * @brief Client for OAuth protected-resource and authorization-server discovery. */ -class OAuthDiscoveryClient { +class MCP_API OAuthDiscoveryClient { public: /** * @brief Construct an OAuth discovery client. @@ -643,8 +485,12 @@ class OAuthDiscoveryClient { */ explicit OAuthDiscoveryClient( std::shared_ptr http_client, - std::chrono::seconds cache_ttl = std::chrono::seconds(constants::g_default_cache_ttl_seconds)) - : http_client_(std::move(http_client)), cache_ttl_(cache_ttl) {} + std::chrono::seconds cache_ttl = std::chrono::seconds(constants::g_default_cache_ttl_seconds)); + + OAuthDiscoveryClient(const OAuthDiscoveryClient&) = delete; + OAuthDiscoveryClient& operator=(const OAuthDiscoveryClient&) = delete; + OAuthDiscoveryClient(OAuthDiscoveryClient&&) = delete; + OAuthDiscoveryClient& operator=(OAuthDiscoveryClient&&) = delete; /** * @brief Discover metadata for a protected resource. @@ -652,47 +498,37 @@ class OAuthDiscoveryClient { * @param resource_url Resource URL whose metadata should be resolved. * @return A task resolving to the discovered protected-resource metadata. */ - Task discover_protected_resource(const std::string& resource_url) { - { - std::lock_guard lock(cache_mutex_); - auto it = resource_cache_.find(resource_url); - if (it != resource_cache_.end() && !it->second.is_expired()) { - co_return it->second.data; - } - } - - auto parsed = parse_url_components(resource_url); - auto base = parsed.scheme + "://" + parsed.authority; - - std::vector urls_to_try; - if (!parsed.path.empty() && parsed.path != "/") { - auto path_part = parsed.path; - if (!path_part.empty() && path_part.front() == '/') { - path_part = path_part.substr(1); - } - urls_to_try.push_back(base + "/.well-known/oauth-protected-resource/" + path_part); - } - urls_to_try.push_back(base + "/.well-known/oauth-protected-resource"); - - for (const auto& url : urls_to_try) { - try { - auto json = co_await http_client_->get_json(url); - auto metadata = json.get(); - - std::lock_guard lock(cache_mutex_); - resource_cache_[resource_url] = { - metadata, - std::chrono::steady_clock::now() + cache_ttl_, - }; - co_return metadata; - } catch (...) { - // Ignore failure and try next fallback URL - continue; - } - } - - throw std::runtime_error("Failed to discover protected resource metadata for " + resource_url); - } + Task discover_protected_resource(const std::string& resource_url); + + /** + * @brief Discover metadata for a protected resource, honouring a challenge-supplied URL. + * + * @param resource_url Resource URL whose metadata should be resolved. + * @param challenge_metadata_url `resource_metadata` URL taken from a `WWW-Authenticate` + * challenge, when the challenge supplied one. + * @return A task resolving to the discovered protected-resource metadata. + * + * @details When the challenge supplied a URL it is fetched and nothing else is tried. Otherwise + * the well-known fallback runs, trying the path-based location before the root one. + */ + Task discover_protected_resource( + const std::string& resource_url, const std::optional& challenge_metadata_url); + + /** + * @brief Discover metadata for a protected resource, subject to the caller's acceptance. + * + * @param resource_url Resource URL whose metadata should be resolved. + * @param challenge_metadata_url `resource_metadata` URL taken from a `WWW-Authenticate` + * challenge, when the challenge supplied one. + * @param accept Called with the document before it is returned or cached; throwing rejects it. + * @return A task resolving to the discovered protected-resource metadata. + * + * @details Identical to the two-argument form except that nothing is written to the cache until + * `accept` has passed on it, and `accept` runs on a cache hit as well. + */ + Task discover_protected_resource( + const std::string& resource_url, const std::optional& challenge_metadata_url, + ProtectedResourceAcceptor accept); /** * @brief Discover metadata for an authorization server. @@ -700,134 +536,31 @@ class OAuthDiscoveryClient { * @param issuer_url Issuer URL or base URL of the authorization server. * @return A task resolving to the discovered authorization-server metadata. */ - Task discover_auth_server(const std::string& issuer_url) { - { - std::lock_guard lock(cache_mutex_); - auto it = auth_cache_.find(issuer_url); - if (it != auth_cache_.end() && !it->second.is_expired()) { - co_return it->second.data; - } - } - - auto parsed = parse_url_components(issuer_url); - auto base = parsed.scheme + "://" + parsed.authority; - - std::vector urls_to_try; - bool has_path = !parsed.path.empty() && parsed.path != "/"; - - if (has_path) { - auto path_part = parsed.path; - if (!path_part.empty() && path_part.front() == '/') { - path_part = path_part.substr(1); - } - if (!path_part.empty() && path_part.back() == '/') { - path_part.pop_back(); - } - urls_to_try.push_back(base + "/.well-known/oauth-authorization-server/" + path_part); - urls_to_try.push_back(base + "/.well-known/openid-configuration/" + path_part); - urls_to_try.push_back(issuer_url + "/.well-known/openid-configuration"); - } else { - urls_to_try.push_back(base + "/.well-known/oauth-authorization-server"); - urls_to_try.push_back(base + "/.well-known/openid-configuration"); - } - - for (const auto& url : urls_to_try) { - try { - auto json = co_await http_client_->get_json(url); - auto metadata = json.get(); - - std::lock_guard lock(cache_mutex_); - auth_cache_[issuer_url] = { - metadata, - std::chrono::steady_clock::now() + cache_ttl_, - }; - co_return metadata; - } catch (...) { - // Ignore failure and try next fallback URL - continue; - } - } - - throw std::runtime_error("Failed to discover authorization server metadata for " + issuer_url); - } + Task discover_auth_server(const std::string& issuer_url); /** * @brief Clear all cached discovery metadata. */ - void clear_cache() { - std::lock_guard lock(cache_mutex_); - resource_cache_.clear(); - auth_cache_.clear(); - } + void clear_cache(); private: - struct UrlComponents { - std::string scheme; - std::string authority; - std::string path; - }; - - static UrlComponents parse_url_components(const std::string& url) { - UrlComponents result; - auto scheme_end = url.find("://"); - if (scheme_end == std::string::npos) { - throw std::invalid_argument("URL missing scheme: " + url); - } - result.scheme = url.substr(0, scheme_end); - auto rest = url.substr(scheme_end + 3); - - auto path_start = rest.find('/'); - if (path_start == std::string::npos) { - result.authority = rest; - result.path = "/"; - } else { - result.authority = rest.substr(0, path_start); - result.path = rest.substr(path_start); - } - return result; - } - - std::shared_ptr http_client_; - std::chrono::seconds cache_ttl_; - - mutable std::mutex cache_mutex_; - std::unordered_map> resource_cache_; - std::unordered_map> auth_cache_; + struct Impl; + std::shared_ptr impl_; }; -/// @brief Callback used to validate a bearer token extracted from request metadata. +/// @brief Legacy callback used to validate a bearer token embedded in request metadata. using TokenValidator = std::function(const std::string& token)>; /** - * @brief Create middleware that validates bearer tokens in request metadata. + * @brief Create legacy middleware that validates bearer tokens in request metadata. + * + * For Streamable HTTP servers, prefer set_bearer_token_validator() on the HTTP transport or + * session manager so authentication is enforced at the HTTP boundary. * * @param validator Async callback that returns true when the token is accepted. * @return Middleware enforcing presence and validity of `_meta.auth_token`. */ -inline Middleware make_auth_middleware(TokenValidator validator) { - return [validator = std::move(validator)](mcp::Context& ctx, const nlohmann::json& params, - TypeErasedHandler next) -> Task { - std::string token; - if (params.contains("_meta") && params["_meta"].contains("auth_token")) { - token = params["_meta"]["auth_token"].get(); - } - - if (token.empty()) { - co_return nlohmann::json{ - {"content", {{{"type", "text"}, {"text", "Unauthorized: missing Bearer token"}}}}, - {"isError", true}}; - } - - bool valid = co_await validator(token); - if (!valid) { - co_return nlohmann::json{ - {"content", {{{"type", "text"}, {"text", "Unauthorized: invalid Bearer token"}}}}, - {"isError", true}}; - } - - co_return co_await next(ctx, params); - }; -} +MCP_API Middleware make_auth_middleware(TokenValidator validator); /** * @brief Extract a bearer token from an Authorization header value. @@ -835,14 +568,7 @@ inline Middleware make_auth_middleware(TokenValidator validator) { * @param auth_header_value Header value to parse. * @return The token value without the `Bearer ` prefix, or an empty string on mismatch. */ -inline std::string extract_bearer_token(std::string_view auth_header_value) { - constexpr std::string_view prefix = "Bearer "; - if (auth_header_value.size() > prefix.size() && - auth_header_value.substr(0, prefix.size()) == prefix) { - return std::string(auth_header_value.substr(prefix.size())); - } - return {}; -} +MCP_API std::string extract_bearer_token(std::string_view auth_header_value); /** * @brief Abstract interface for providing access tokens and handling refresh. @@ -869,71 +595,213 @@ class Authenticator { * @return true if refresh succeeded and a new token was persisted; false otherwise. */ virtual Task try_refresh_token() = 0; + + /** + * @brief Performs challenge-driven authorization for a `WWW-Authenticate` response. + * + * @param www_authenticate Raw `WWW-Authenticate` header value from the challenge response. + * @return true if authorization completed and a new token was persisted; false otherwise. + * + * @details Called before try_refresh_token() when a challenge is available, so an implementation + * can discover metadata and run a full authorization exchange rather than only renewing an + * existing grant. The default implementation reports that it handled nothing. + */ + virtual Task try_handle_challenge(const std::string& www_authenticate) { + (void)www_authenticate; + co_return false; + } + + /** + * @brief Cancel any authorization work this authenticator has in flight and release parked + * callers with an error. + * + * @details Called once, synchronously, by the owning transport's close(). An implementation that + * holds no cancellable network state may leave the default no-op. Safe to call more than once. + */ + virtual void close() {} }; /** * @brief OAuth 2.0 implementation of the Authenticator interface. */ -class OAuthAuthenticator : public Authenticator { +class MCP_API OAuthAuthenticator : public Authenticator { public: OAuthAuthenticator(std::shared_ptr token_store, std::shared_ptr oauth_client, OAuthConfig config, - std::string server_url) - : token_store_(std::move(token_store)), - oauth_client_(std::move(oauth_client)), - config_(std::move(config)), - server_url_(std::move(server_url)) {} - - [[nodiscard]] std::string get_access_token() const override { - auto token = token_store_->load(server_url_); - if (token) { - return token->access_token; - } - return {}; - } + std::string server_url); - Task try_refresh_token() override { - auto stored = token_store_->load(server_url_); - if (!stored || !stored->refresh_token) { - co_return false; - } - - try { - auto new_token = co_await oauth_client_->refresh_token(config_, *stored->refresh_token); - if (!new_token.refresh_token && stored->refresh_token) { - new_token.refresh_token = stored->refresh_token; - } - token_store_->store(server_url_, std::move(new_token)); - co_return true; - } catch (...) { - co_return false; - } - } + [[nodiscard]] std::string get_access_token() const override; + + Task try_refresh_token() override; /// @brief Stores an initial token obtained from an explicit OAuth exchange (e.g., authorization /// code flow). /// @param token The token response to persist via the configured TokenStore. - void store_token(TokenResponse token) { token_store_->store(server_url_, std::move(token)); } + void store_token(TokenResponse token); + + /// @brief Aborts any in-flight token-refresh HTTP exchange. + void close() override; + + private: + struct Impl; + std::shared_ptr impl_; +}; + +/** + * @brief Application-controlled consent step for an authorization attempt. + * + * @details Receives the per-attempt request record and returns the response the authorization + * server delivered to the redirect URI. The SDK never launches a browser and never binds an + * unsolicited listener; carrying the user agent to the authorization endpoint and collecting the + * redirect is entirely the application's responsibility. + */ +using AuthorizationCallback = std::function(const AuthorizationRequest&)>; + +/** + * @brief Configuration for challenge-driven OAuth authorization. + */ +struct OAuthAuthorizationConfig { + std::string server_url; ///< MCP server URL that issued the challenge; also the token store key. + /// Client identifier presented to the authorization server. Setting it is shorthand for + /// injecting pre-registered credentials: it is treated exactly like + /// `client_identity.pre_registered` and therefore never falls back to registration. + std::string client_id; + std::optional client_secret; ///< Optional confidential-client secret. + /// Issuer the shorthand credentials above are bound to. Required whenever `client_secret` is + /// set: the authorization server is named by the protected-resource document, so a secret that + /// names no issuer is refused rather than presented. "Bound to no issuer" must never be read as + /// "bound to every issuer". Leave empty for a public client, whose `client_id` is not a secret + /// and may be presented to any authorization server. + std::string client_issuer; + std::string redirect_uri; ///< Redirect URI the authorization response returns to. + /// Scope override. When set it wins over both the challenge scope and the resource metadata; + /// when unset the challenge scope is preferred, then `scopes_supported`, then no scope at all. + std::optional scope; + /// Client identity inputs consulted when no `client_id` was supplied: a published client ID + /// metadata document URL, injected credentials, and the metadata used for registration. + ClientIdentityConfig client_identity; + /// Issuer-keyed storage for credentials obtained by dynamic registration. When unset, a + /// registration is performed per authorization attempt rather than reused. + std::shared_ptr credential_store; + MetadataFetchPolicy policy; ///< Outbound-request policy for every discovery and token request. + HostResolver host_resolver; ///< Optional custom resolver; the system resolver is used when unset. +}; + +/** + * @brief Authenticator that performs challenge-driven OAuth authorization. + * + * @details Composes the pieces a `WWW-Authenticate` response requires: challenge parsing, + * protected-resource and authorization-server discovery under the configured fetch policy, an + * authorization request carrying S256 PKCE and cryptographic `state` bound to the issuer recorded + * from the selected metadata document, RFC 9207 response validation, and an authorization-code + * exchange carrying the RFC 8707 `resource` indicator. + */ +class MCP_API OAuthAuthorizationManager : public Authenticator { + public: + /** + * @brief Construct a challenge-driven authorization manager. + * + * @param executor Executor used for asynchronous operations. + * @param token_store Storage for the acquired access token. + * @param config Client identity, redirect URI and outbound-request policy. + * @param callback Application consent step invoked once per authorization attempt. + */ + OAuthAuthorizationManager(const net::any_io_executor& executor, + std::shared_ptr token_store, OAuthAuthorizationConfig config, + AuthorizationCallback callback); + + /** + * @brief Return the stored access token without network I/O. + * + * @return The stored access token, or an empty string when none has been acquired. + */ + [[nodiscard]] std::string get_access_token() const override; + + /** + * @brief Renew the stored token using its refresh token, when one is present. + * + * @return true when a renewed token was persisted. + */ + Task try_refresh_token() override; + + /** + * @brief Run a full authorization exchange for a challenge response. + * + * @param www_authenticate Raw `WWW-Authenticate` header value. + * @return true when authorization completed and an access token was persisted; false when the + * response carried no `Bearer` challenge to act on. + * + * @throws MetadataPolicyError If any discovery or token target is refused by the fetch policy. + * @throws std::runtime_error If discovery fails or the authorization response is rejected. + */ + Task try_handle_challenge(const std::string& www_authenticate) override; + + /** + * @brief Return the record of the most recent authorization attempt. + * + * @return The request record, or `std::nullopt` when no attempt has been made. + * + * @details Exposes the state, PKCE verifier, recorded issuer and resource indicator that were + * actually used, so applications can audit the binding an attempt was validated against. + */ + [[nodiscard]] std::optional last_authorization_request() const; + + /** + * @brief Return the client identity used by the most recent authorization attempt. + * + * @return The identity, or `std::nullopt` when no attempt has resolved one. + * + * @details `source` records which path produced it, so an application can tell an injected + * credential from a metadata-document identifier from a dynamic registration. + */ + [[nodiscard]] std::optional last_client_identity() const; + + /** + * @brief Abort this manager's HTTP work, release parked followers with an error, and refuse + * further authorization attempts. + * + * @details Aborts the in-progress discovery or token-exchange HTTP exchange, wakes every + * coalesced follower, and refuses new authorization attempts from this point on -- including one + * racing this very call. Safe to call from any thread. Idempotent. + * + * @note A leader parked inside the application's own authorization callback is not reachable from + * here. Followers coalesced onto that flow are still released, but the flow itself ends only when + * the application's callback returns or its executor stops. + */ + void close() override; private: - std::shared_ptr token_store_; - std::shared_ptr oauth_client_; - OAuthConfig config_; - std::string server_url_; + struct Impl; + std::shared_ptr impl_; +}; + +/** + * @brief Resource limits for legacy JSON-RPC authentication replay correlation. + * + * @details Pending request wires are retained only until a matching response arrives, the entry is + * evicted to satisfy these limits, the TTL elapses, or the transport closes. A zero limit or a + * non-positive TTL disables legacy response-driven replay correlation. HTTP status-driven refresh + * and retry does not depend on this cache. + */ +struct OAuthClientTransportOptions { + std::size_t max_pending_requests{256}; ///< Maximum retained request count. + std::size_t max_pending_request_bytes{16 * 1024 * 1024}; ///< Approximate retained wire/key bytes. + std::chrono::milliseconds pending_request_ttl{std::chrono::minutes(5)}; ///< Replay eligibility. }; /** - * @brief Transport wrapper that injects and refreshes OAuth bearer tokens. + * @brief Transport wrapper that supplies and refreshes OAuth bearer tokens. * - * @details Wraps any `ITransport` to automatically inject OAuth bearer tokens on write and handle token - * refresh on auth failures. When the server returns error codes -32001 or -32000 - * (authentication-related), the transport calls `Authenticator::try_refresh_token()` and, if - * successful, re-sends the last written message and reads the new response. Only one retry attempt is + * @details For HttpClientTransport, tokens are sent in the HTTP Authorization header. The legacy + * request-metadata mechanism is retained only for non-HTTP transports. The wrapper also handles token + * refresh on legacy JSON-RPC authorization failures. When the server returns -32000, the + * transport calls `Authenticator::try_refresh_token()` and, if successful, re-sends the request whose + * ID matches the error response and reads the new response. Only one retry attempt is * made per `read_message()` call. Callers should ensure messages are idempotent since they may be - * re-sent after a token refresh. `last_written_message_` stores the most recently written message for - * potential replay. + * re-sent after a token refresh. Outstanding request wires are tracked by JSON-RPC request ID in a + * bounded, expiring cache so an authentication error cannot replay a different concurrent request. */ -class OAuthClientTransport final : public ITransport { +class MCP_API OAuthClientTransport final : public ITransport { public: /** * @brief Construct an authenticated transport wrapper. @@ -942,49 +810,30 @@ class OAuthClientTransport final : public ITransport { * @param authenticator Authenticator used to retrieve and refresh tokens. */ OAuthClientTransport(std::shared_ptr inner, - std::shared_ptr authenticator) - : inner_(std::move(inner)), authenticator_(std::move(authenticator)) {} + std::shared_ptr authenticator); + + /** + * @brief Construct an authenticated transport wrapper with replay-correlation limits. + * + * @param inner Underlying transport used for MCP message exchange. + * @param authenticator Authenticator used to retrieve and refresh tokens. + * @param options Bounds and TTL for retaining outstanding request wires. + */ + OAuthClientTransport(std::shared_ptr inner, + std::shared_ptr authenticator, + OAuthClientTransportOptions options); /** * @brief Read a message from the inner transport, retrying once on authentication errors. - * @details If the received message contains an authentication error (JSON-RPC error code -32001 or - * -32000), attempts one token refresh via `Authenticator::try_refresh_token()`. On successful - * refresh, replays the last written message and returns the new response. If refresh fails or a - * second auth error is received, returns the error response as-is. If the response cannot be parsed - * as JSON, the raw string is returned unchanged. + * @details If the received message contains the legacy JSON-RPC authentication error -32000, + * attempts one token refresh via `Authenticator::try_refresh_token()`. On successful + * refresh, replays the matching outstanding request while it remains eligible and returns the new + * response. If refresh fails, correlation has expired or been evicted, or a second auth error is + * received, returns the error response as-is. If the response cannot be parsed as JSON, the raw + * string is returned unchanged. * @return The (possibly retried) raw message string. */ - Task read_message() override { - if (!inner_) { - throw std::runtime_error("OAuthClientTransport inner transport is null"); - } - auto raw = co_await inner_->read_message(); - - try { - auto json_msg = nlohmann::json::parse(raw); - JSONRPCMessage msg = json_msg.get(); - - if (auto* error_resp = std::get_if(&msg)) { - auto code = error_resp->error.code; - // Error codes -32001 (authentication required) and -32000 (authentication failed) - // are MCP standard JSON-RPC error codes for auth-related failures. - // Attempt to refresh token and retry the original request once. - if (code == g_REQUEST_TIMEOUT || code == g_UNAUTHORIZED) { - bool refreshed = co_await authenticator_->try_refresh_token(); - if (refreshed && !last_written_message_.empty()) { - co_await write_message(last_written_message_); - co_return co_await inner_->read_message(); - } - } - } - } catch (const std::exception& e) { - // Silently ignore JSON parse errors: returns raw message to caller unchanged. - // This allows the client to handle non-JSON responses gracefully. - (void)e; - } - - co_return raw; - } + Task read_message() override; /** * @brief Inject the current bearer token into an outgoing MCP message. @@ -992,59 +841,16 @@ class OAuthClientTransport final : public ITransport { * @param message Serialized JSON-RPC request or notification. * @return A task that completes once the wrapped transport accepts the message. */ - Task write_message(std::string_view message) override { - if (!inner_) { - throw std::runtime_error("OAuthClientTransport inner transport is null"); - } - std::string injected = std::string(message); - try { - auto json_msg = nlohmann::json::parse(message); - JSONRPCMessage msg = json_msg.get(); - - if (std::holds_alternative(msg)) { - injected = inject_token(std::get(msg)); - } - } catch (const std::exception& e) { - // Ignore parse errors, send original message. - (void)e; - } - - last_written_message_ = std::string(message); - co_await inner_->write_message(injected); - } + Task write_message(std::string_view message) override; /** * @brief Close the wrapped transport. */ - void close() override { inner_->close(); } + void close() override; private: - [[nodiscard]] std::string inject_token(JSONRPCRequest request) const { - auto token = authenticator_->get_access_token(); - if (token.empty()) { - return nlohmann::json(request).dump(); - } - - if (!request.params) { - request.params = nlohmann::json::object(); - } - - auto& params = *request.params; - if (!params.is_object()) { - params = nlohmann::json::object(); - } - - if (!params.contains("_meta")) { - params["_meta"] = nlohmann::json::object(); - } - params["_meta"]["auth_token"] = token; - - return nlohmann::json(request).dump(); - } - - std::shared_ptr inner_; - std::shared_ptr authenticator_; - std::string last_written_message_; + struct Impl; + std::shared_ptr impl_; }; } // namespace mcp::auth diff --git a/include/mcp/client/client.hpp b/include/mcp/client/client.hpp index 19d87b5..c8d35e5 100644 --- a/include/mcp/client/client.hpp +++ b/include/mcp/client/client.hpp @@ -13,20 +13,16 @@ #include #include +#include #include #include -#include #include -#include -#include -#include -#include -#include -#include +#include #include #include #include +#include #include #include #include @@ -35,6 +31,80 @@ namespace mcp { +/** + * @brief Structured error returned by an MCP peer or generated by the client runtime. + * + * Unlike a plain std::runtime_error, this preserves the JSON-RPC error code and optional + * data payload so callers can make protocol-aware decisions. + */ +class MCP_API McpError : public std::runtime_error { + public: + /** @brief Construct from a protocol error object. */ + explicit McpError(Error error); + + /** + * @brief Construct a protocol-aware error. + * @param code JSON-RPC or SDK-defined error code. + * @param message Human-readable error message. + * @param data Optional structured error details. + */ + McpError(int code, std::string message, std::optional data = std::nullopt); + + /** @brief Return the JSON-RPC or SDK-defined error code. */ + [[nodiscard]] int code() const noexcept { return error_.code; } + /** @brief Return the protocol error message. */ + [[nodiscard]] const std::string& message() const noexcept { return error_.message; } + /** @brief Return optional structured error details. */ + [[nodiscard]] const std::optional& data() const noexcept { return error_.data; } + /** @brief Return the complete protocol error object. */ + [[nodiscard]] const Error& error() const noexcept { return error_; } + + private: + Error error_; +}; + +/** @brief Client-wide runtime behavior. */ +struct ClientOptions { + /** + * @brief Default time limit for outgoing requests. Thirty seconds unless set. + * + * Every request the client issues carries this deadline, including the typed helpers such + * as call_tool() and read_resource(), which accept no per-request override. A request that + * outlives it fails with McpError g_REQUEST_TIMEOUT while the peer may still be running it, + * so raise this for a client whose tools legitimately take longer. Must be positive: the + * constructor throws std::invalid_argument otherwise, and there is no value that disables + * the deadline. + */ + std::chrono::milliseconds request_timeout{std::chrono::seconds(30)}; + /** @brief Reject malformed JSON-RPC envelopes received from the peer. */ + bool strict_protocol_validation{true}; + /** + * @brief Invoked when an incoming message is discarded instead of dispatched. + * + * Reports two things, neither of which ends the session and neither of which the + * application can otherwise observe: + * - a message the peer sent that the client could not decode, as g_PARSE_ERROR; + * - an exception thrown by an application notification callback, as g_INTERNAL_ERROR, + * with the notification method in the message. + * + * Runs on the client's read loop, so it must not block. An exception thrown from it is + * discarded: there is nowhere left to report a failure of the reporting path itself. + */ + std::function on_protocol_error; +}; + +/** @brief Overrides for one outgoing request. */ +struct RequestOptions { + /** + * @brief Per-request deadline; uses ClientOptions::request_timeout when absent. + * + * A write still queued at the deadline is canceled before transmission. If + * transport I/O has already started, remote execution may be ambiguous even + * though the caller receives a timeout, as with any network request deadline. + */ + std::optional timeout; +}; + /** * @brief MCP client that communicates with a server over a transport. * @@ -44,10 +114,9 @@ namespace mcp { * via connect(), which sends an initialize request followed by an * initialized notification. * - * All async operations execute on a boost::asio::strand for thread safety - * without std::mutex. + * Runtime I/O and request state are serialized on a boost::asio::strand. */ -class Client { +class MCP_API Client { public: /** * @brief Construct a Client. @@ -55,54 +124,42 @@ class Client { * @param transport The transport to use for message exchange. * Ownership is shared with the client. * @param executor The executor to use for async operations. + * @param options Default timeout and protocol-validation behavior. */ - Client(std::shared_ptr transport, const boost::asio::any_io_executor& executor) - : transport_(std::move(transport)), strand_(boost::asio::make_strand(executor)) {} + Client(std::shared_ptr transport, const boost::asio::any_io_executor& executor, + ClientOptions options = {}); /// @brief Destructor. Closes the transport and clears pending requests /// so that asio-tied objects (timers, streams) are destroyed while the /// io_context is still alive. - ~Client() { - if (transport_) { - transport_->close(); - } - pending_requests_.clear(); - } + ~Client(); Client(const Client&) = delete; Client& operator=(const Client&) = delete; Client(Client&&) = delete; Client& operator=(Client&&) = delete; + /** + * @brief Perform the MCP initialization handshake with default capabilities. + * @param name Client implementation name. + * @param version Client implementation version. + * @return A task that resolves to the server's InitializeResult. + */ + // NOLINTNEXTLINE(bugprone-easily-swappable-parameters) + Task connect(std::string_view name, std::string_view version); + /** * @brief Perform the MCP initialization handshake. * - * @details Starts the read loop via co_spawn, sends an initialize - * request with the given client info and capabilities, waits for the - * server response, then sends an initialized notification. + * @details Starts the read loop, sends initialize, waits for the server + * response, and then sends notifications/initialized. * * @param client_info Information about this client implementation. - * @param capabilities The capabilities this client supports. + * @param capabilities Capabilities supported by this client. * @return A task that resolves to the server's InitializeResult. */ - // NOLINTNEXTLINE(bugprone-easily-swappable-parameters) - Task connect(std::string_view name, std::string_view version) { - Implementation info; - info.name = std::string(name); - info.version = std::string(version); - return connect(info, {}); - } - Task connect(const Implementation& client_info, - const ClientCapabilities& capabilities) { - // [gcc11-sso: not-a-coroutine] DO NOT add co_await / co_return here. - boost::asio::co_spawn(strand_, read_loop(), boost::asio::detached); - InitializeRequest init_req; - init_req.protocolVersion = std::string(g_LATEST_PROTOCOL_VERSION); - init_req.clientInfo = client_info; - init_req.capabilities = capabilities; - return connect_impl(nlohmann::json(std::move(init_req))); - } + const ClientCapabilities& capabilities); /** * @brief Send a JSON-RPC request and await its response. @@ -114,26 +171,18 @@ class Client { * @param method JSON-RPC method name. * @param params Optional parameters for the request. * @return A task that resolves to the result JSON object. - * @throws std::runtime_error If the server returns an error response. + * @throws McpError If the peer returns an error response or the request times out. */ // [gcc11-sso: string_view] Task send_request(std::string_view method, - const std::optional& params) { - int64_t id = next_request_id_.fetch_add(1, std::memory_order_relaxed); - auto id_str = std::to_string(id); + const std::optional& params); - JSONRPCRequest request; - request.id = RequestId{id_str}; - request.method = std::string(method); - request.params = params; - - auto& pending = pending_requests_[id_str]; - pending.timer = std::make_unique(strand_); - pending.timer->expires_at(std::chrono::steady_clock::time_point::max()); - auto* timer_ptr = pending.timer.get(); - - return send_request_impl(nlohmann::json(std::move(request)).dump(), timer_ptr, id); - } + /** + * @brief Send a JSON-RPC request with per-request runtime options. + */ + Task send_request(std::string_view method, + const std::optional& params, + const RequestOptions& request_options); /** * @brief Send a JSON-RPC notification (no response expected). @@ -143,27 +192,17 @@ class Client { * @return A task that completes when the notification has been written to the transport. */ // [gcc11-sso: string_view] - Task send_notification(std::string_view method, const std::optional& params) { - JSONRPCNotification notification; - notification.method = std::string(method); - notification.params = params; - return send_notification_impl(nlohmann::json(std::move(notification)).dump()); - } + Task send_notification(std::string_view method, const std::optional& params); /** * @brief Get the number of pending (in-flight) requests. * * @return The count of requests awaiting responses. */ - [[nodiscard]] std::size_t pending_request_count() const { return pending_requests_.size(); } + [[nodiscard]] std::size_t pending_request_count() const; /// @brief Close the client, stopping the read loop and releasing resources. - void close() { - if (transport_) { - transport_->close(); - } - pending_requests_.clear(); - } + void close(); /** * @brief Call a tool on the server. @@ -191,16 +230,7 @@ class Client { * @param cursor Optional pagination cursor. * @return A task that resolves to the server's ListToolsResult. */ - Task list_tools(const std::optional& cursor = std::nullopt) { - // [gcc11-sso: not-a-coroutine] - std::optional params; - if (cursor) { - PaginatedRequestParams paginated; - paginated.cursor = cursor; - params = nlohmann::json(std::move(paginated)); - } - return call_and_parse("tools/list", std::move(params)); - } + Task list_tools(const std::optional& cursor = std::nullopt); /** * @brief List available resources from the server. @@ -208,16 +238,7 @@ class Client { * @param cursor Optional pagination cursor. * @return A task that resolves to the server's ListResourcesResult. */ - Task list_resources(const std::optional& cursor = std::nullopt) { - // [gcc11-sso: not-a-coroutine] - std::optional params; - if (cursor) { - PaginatedRequestParams paginated; - paginated.cursor = cursor; - params = nlohmann::json(std::move(paginated)); - } - return call_and_parse("resources/list", std::move(params)); - } + Task list_resources(const std::optional& cursor = std::nullopt); /** * @brief Read a resource from the server by URI. @@ -225,12 +246,7 @@ class Client { * @param uri The URI of the resource to read. * @return A task that resolves to the server's ReadResourceResult. */ - Task read_resource(const std::string& uri) { - // [gcc11-sso: not-a-coroutine] - ReadResourceRequestParams p; - p.uri = uri; - return call_and_parse("resources/read", nlohmann::json(std::move(p))); - } + Task read_resource(const std::string& uri); /** * @brief List available resource templates from the server. @@ -239,17 +255,7 @@ class Client { * @return A task that resolves to the server's ListResourceTemplatesResult. */ Task list_resource_templates( - const std::optional& cursor = std::nullopt) { - // [gcc11-sso: not-a-coroutine] - std::optional params; - if (cursor) { - PaginatedRequestParams paginated; - paginated.cursor = cursor; - params = nlohmann::json(std::move(paginated)); - } - return call_and_parse("resources/templates/list", - std::move(params)); - } + const std::optional& cursor = std::nullopt); /** * @brief List available prompts from the server. @@ -257,16 +263,7 @@ class Client { * @param cursor Optional pagination cursor. * @return A task that resolves to the server's ListPromptsResult. */ - Task list_prompts(const std::optional& cursor = std::nullopt) { - // [gcc11-sso: not-a-coroutine] - std::optional params; - if (cursor) { - PaginatedRequestParams paginated; - paginated.cursor = cursor; - params = nlohmann::json(std::move(paginated)); - } - return call_and_parse("prompts/list", std::move(params)); - } + Task list_prompts(const std::optional& cursor = std::nullopt); /** * @brief Get a specific prompt from the server. @@ -277,13 +274,7 @@ class Client { */ Task get_prompt( const std::string& name, - const std::optional>& arguments = std::nullopt) { - // [gcc11-sso: not-a-coroutine] - GetPromptRequestParams p; - p.name = name; - p.arguments = arguments; - return call_and_parse("prompts/get", nlohmann::json(std::move(p))); - } + const std::optional>& arguments = std::nullopt); /** * @brief Request completions from the server. @@ -291,17 +282,14 @@ class Client { * @param params The completion parameters. * @return A task that resolves to the server's CompleteResult. */ - Task complete(const CompleteParams& params) { - // [gcc11-sso: not-a-coroutine] - return call_and_parse("completion/complete", nlohmann::json(params)); - } + Task complete(const CompleteParams& params); /** * @brief Send a ping request to the connected server. * * @return A task that completes once the ping round-trip succeeds. */ - Task ping() { co_await send_request("ping", std::nullopt); } + Task ping(); /** * @brief Send a cancellation notification for an in-flight request. @@ -311,13 +299,7 @@ class Client { * @return A task that completes when the cancellation notification has been sent. */ Task cancel(const RequestId& request_id, - const std::optional& reason = std::nullopt) { - // [gcc11-sso: not-a-coroutine] - CancelledNotificationParams params; - params.requestId = request_id; - params.reason = reason; - return send_notification("notifications/cancelled", nlohmann::json(std::move(params))); - } + const std::optional& reason = std::nullopt); /// @brief Type for notification callback: receives the notification params JSON. using NotificationCallback = std::function; @@ -334,9 +316,7 @@ class Client { * @param method The notification method to listen for (e.g. "notifications/progress"). * @param callback The callback to invoke when the notification arrives. */ - void on_notification(const std::string& method, NotificationCallback callback) { - notification_handlers_[method] = std::move(callback); - } + void on_notification(const std::string& method, NotificationCallback callback); /// @brief Callback invoked for a deserialized progress notification. using ProgressCallback = std::function; @@ -346,12 +326,7 @@ class Client { * * @param callback The callback invoked when a notifications/progress message arrives. */ - void on_progress(ProgressCallback callback) { - on_notification("notifications/progress", - [cb = std::move(callback)](const nlohmann::json& params) { - cb(params.get()); - }); - } + void on_progress(ProgressCallback callback); /** * @brief Register a handler for incoming requests from the server (reverse RPC). @@ -363,9 +338,7 @@ class Client { * @param method The request method to handle. * @param handler The async handler that returns a result JSON. */ - void on_request(const std::string& method, RequestHandler handler) { - request_handlers_[method] = std::move(handler); - } + void on_request(const std::string& method, RequestHandler handler); /** * @brief Register a handler for incoming elicitation requests from the server. @@ -375,15 +348,7 @@ class Client { * * @param handler Async handler: receives ElicitRequestParams, returns ElicitResult. */ - void on_elicitation(std::function(const ElicitRequestParams&)> handler) { - on_request("elicitation/create", - [h = std::move(handler)](const nlohmann::json& params) -> Task { - auto elicit_params = params.get(); - auto result = co_await h(elicit_params); - nlohmann::json result_json = std::move(result); - co_return result_json; - }); - } + void on_elicitation(std::function(const ElicitRequestParams&)> handler); /** * @brief Set static roots and auto-register the roots/list handler. @@ -395,249 +360,30 @@ class Client { * @param roots The list of roots to provide. * @param notify Whether to send a roots-changed notification. */ - void set_roots(const std::vector& roots, bool notify = false) { - roots_ = roots; - - if (request_handlers_.find("roots/list") == request_handlers_.end()) { - on_request("roots/list", [this](const nlohmann::json&) -> Task { - ListRootsResult result; - result.roots = roots_; - nlohmann::json j = std::move(result); - co_return j; - }); - } - - bool server_supports_roots_changed = - server_capabilities_.resources.has_value() && - server_capabilities_.resources->listChanged.value_or(false); - - if (notify && server_supports_roots_changed && transport_) { - boost::asio::co_spawn( - strand_, - [this]() -> Task { - co_await send_notification("notifications/roots/list_changed", std::nullopt); - }, - boost::asio::detached); - } - } + void set_roots(const std::vector& roots, bool notify = false); /** * @brief Register a dynamic handler for roots/list requests. * * @param handler Async handler that returns a ListRootsResult. */ - void on_roots_list(std::function(const nlohmann::json&)> handler) { - on_request("roots/list", - [h = std::move(handler)](const nlohmann::json& params) -> Task { - auto result = co_await h(params); - nlohmann::json j = std::move(result); - co_return j; - }); - } + void on_roots_list(std::function(const nlohmann::json&)> handler); private: - struct PendingRequest { - std::unique_ptr timer; - nlohmann::json result; - std::optional error; - }; - - // [gcc11-sso: not-a-coroutine] DO NOT change init_params to a type containing std::string. - Task connect_impl(nlohmann::json init_params) { - auto result_json = co_await send_request("initialize", std::move(init_params)); - server_capabilities_ = result_json.get().capabilities; - co_await send_notification("notifications/initialized", std::nullopt); - co_return result_json.get(); - } - - Task send_request_impl(std::string wire, boost::asio::steady_timer* timer_ptr, - int64_t id) { - co_await transport_->write_message(wire); - - try { - co_await timer_ptr->async_wait(boost::asio::use_awaitable); - } catch (const boost::system::system_error& err) { - if (err.code() != boost::asio::error::operation_aborted) { - throw; - } - } - - auto id_str = std::to_string(id); - auto it = pending_requests_.find(id_str); - if (it == pending_requests_.end()) { - throw std::runtime_error("pending request not found for id: " + id_str); - } - auto result = std::move(it->second.result); - auto error = std::move(it->second.error); - pending_requests_.erase(it); - - if (error) { - throw std::runtime_error("JSON-RPC error " + std::to_string(error->code) + ": " + - error->message); - } - co_return result; - } - - Task send_notification_impl(std::string wire) { co_await transport_->write_message(wire); } - // [gcc11-sso: string_view] nlohmann::json is heap-allocated; string_view is trivially copyable. template Task call_and_parse(std::string_view method, std::optional params) { - auto result_json = co_await send_request(method, params); - co_return result_json.template get(); + return parse_result(send_request(method, params)); } - /** - * @brief Infinite read loop that dispatches incoming messages. - * - * @details Continuously reads messages from the transport, classifies - * them as responses, requests, or notifications, and dispatches each - * to the appropriate handler. Responses wake their pending coroutine. - * Requests are dispatched to registered request handlers and their - * result is sent back. Notifications invoke registered callbacks. - */ - Task read_loop() { - if (!transport_) { - co_return; - } - try { - for (;;) { - auto raw = co_await transport_->read_message(); - auto json_msg = nlohmann::json::parse(raw); - - bool has_id = json_msg.contains("id"); - bool has_method = json_msg.contains("method"); - - if (has_id && !has_method) { - dispatch_response(json_msg); - } else if (has_id && has_method) { - boost::asio::co_spawn(strand_, dispatch_incoming_request(std::move(json_msg)), - boost::asio::detached); - } else if (!has_id && has_method) { - dispatch_notification(json_msg); - } - } - } catch (const std::exception& e) { - // Transport closed or error occurred, terminate loop. - (void)e; - } - } - - void dispatch_response(const nlohmann::json& json_msg) { - auto id = json_msg.at("id").get(); - auto id_str = id.to_string(); - - auto it = pending_requests_.find(id_str); - if (it == pending_requests_.end()) { - return; - } - - if (json_msg.contains("error")) { - it->second.error = json_msg.at("error").get(); - } else if (json_msg.contains("result")) { - it->second.result = json_msg.at("result"); - } - - it->second.timer->cancel(); - } - - void dispatch_notification(const nlohmann::json& json_msg) { - auto method = json_msg.at("method").get(); - - if (method == "notifications/cancelled") { - if (json_msg.contains("params")) { - auto params = json_msg.at("params").get(); - auto id_str = params.requestId.to_string(); - - auto it = pending_requests_.find(id_str); - if (it != pending_requests_.end()) { - it->second.error = Error{g_REQUEST_CANCELLED, "Request cancelled by server"}; - it->second.timer->cancel(); - } - } - return; - } - - auto it = notification_handlers_.find(method); - if (it != notification_handlers_.end()) { - auto params = json_msg.contains("params") ? json_msg.at("params") : nlohmann::json{}; - it->second(params); - } - } - - /** - * @brief Handles an incoming JSON-RPC request from the server (reverse RPC). - * - * @details Looks up the method in request_handlers_ and invokes the registered - * handler. Exceptions from handlers are caught and returned as JSON-RPC error - * responses with code -32603 (g_INTERNAL_ERROR). The @c ping method is handled - * automatically and returns an empty result object. Unknown methods return a - * g_METHOD_NOT_FOUND error response. Runs on the client's asio strand - * (co_spawned from the read loop). - * - * @param json_msg The raw JSON-RPC request with @c id, @c method, and optional @c params. - */ - Task dispatch_incoming_request(nlohmann::json json_msg) { - // [gcc11-sso: string_view] Points into heap-allocated json_msg. DO NOT change to std::string. - std::string_view method = json_msg.at("method").get_ref(); - - auto it = request_handlers_.find(method); - if (it != request_handlers_.end()) { - auto params = json_msg.contains("params") ? json_msg.at("params") : nlohmann::json{}; - // [gcc11-sso: scope-before-await] nlohmann::json is heap-safe. DO NOT use optional. - nlohmann::json error_payload; - nlohmann::json result; - try { - result = co_await it->second(params); - } catch (const std::exception& e) { - error_payload = e.what(); - } - // Re-extract id just-in-time AFTER suspension — never stored in frame as RequestId. - if (!error_payload.is_null()) { - co_await transport_->write_message(make_error_wire(json_msg.at("id").get(), - g_INTERNAL_ERROR, - error_payload.get())); - } else { - co_await transport_->write_message( - make_result_wire(json_msg.at("id").get(), std::move(result))); - } - } else if (method == "ping") { - co_await transport_->write_message( - make_result_wire(json_msg.at("id").get(), nlohmann::json::object())); - } else { - co_await transport_->write_message( - make_error_wire(json_msg.at("id").get(), g_METHOD_NOT_FOUND, - "Method not found: " + std::string(method))); - } - } - - // [gcc11-sso: wire-builders] DO NOT convert to Task. - static std::string make_result_wire(const RequestId& id, nlohmann::json result) { - JSONRPCResultResponse response; - response.id = id; - response.result = std::move(result); - return nlohmann::json(std::move(response)).dump(); - } - - static std::string make_error_wire(const RequestId& id, int code, std::string message) { - Error error; - error.code = code; - error.message = std::move(message); - JSONRPCErrorResponse response; - response.id = id; - response.error = std::move(error); - return nlohmann::json(std::move(response)).dump(); + template + static Task parse_result(Task request) { + auto result_json = co_await std::move(request); + co_return result_json.template get(); } - std::shared_ptr transport_; - boost::asio::strand strand_; - std::map pending_requests_; - std::atomic next_request_id_{1}; - - std::map> notification_handlers_; - std::map> request_handlers_; - std::vector roots_; - ServerCapabilities server_capabilities_; + struct Impl; + std::shared_ptr impl_; }; } // namespace mcp diff --git a/include/mcp/core/context.hpp b/include/mcp/core/context.hpp index 43dbadc..e320c21 100644 --- a/include/mcp/core/context.hpp +++ b/include/mcp/core/context.hpp @@ -22,6 +22,9 @@ namespace mcp { /// @brief Type alias for a function that sends a JSON-RPC request and returns the result. using RequestSender = std::function(std::string, std::optional)>; +/// @brief Type alias for a function that serializes outbound notification writes. +using MessageSender = std::function(std::string_view)>; + /** * @brief Request-scoped context provided to server handlers. * @@ -59,14 +62,18 @@ class Context { * @param cancelled Shared cancellation flag set by the server on notifications/cancelled. * @param progress_token Optional progress token from the request's _meta. * @param log_level Pointer to the server's current log level (shared across all contexts). + * @param message_sender Optional serialized notification sender. When omitted, Context writes + * directly to the transport. */ Context(ITransport& transport, RequestSender sender, std::shared_ptr> cancelled, - std::optional progress_token, const std::atomic* log_level) + std::optional progress_token, const std::atomic* log_level, + MessageSender message_sender = {}) : transport_(transport), sender_(std::move(sender)), cancelled_(std::move(cancelled)), progress_token_(std::move(progress_token)), - log_level_(log_level) {} + log_level_(log_level), + message_sender_(std::move(message_sender)) {} /** * @brief Sends an informational log message to the client. @@ -99,7 +106,7 @@ class Context { } // [gcc11-sso: scope-before-await] - std::string wire; + std::shared_ptr wire; { LoggingMessageNotificationParams params; params.level = level; @@ -109,9 +116,13 @@ class Context { JSONRPCNotification notification; notification.method = "notifications/message"; notification.params = nlohmann::json(std::move(params)); - wire = nlohmann::json(std::move(notification)).dump(); + wire = std::make_shared(nlohmann::json(std::move(notification)).dump()); + } + if (message_sender_) { + co_await message_sender_(*wire); + } else { + co_await transport_.write_message(*wire); } - co_await transport_.write_message(wire); } /** @@ -145,7 +156,7 @@ class Context { } // [gcc11-sso: scope-before-await] - std::string wire; + std::shared_ptr wire; { ProgressNotificationParams params; params.progressToken = *progress_token_; @@ -156,9 +167,13 @@ class Context { JSONRPCNotification notification; notification.method = "notifications/progress"; notification.params = nlohmann::json(std::move(params)); - wire = nlohmann::json(std::move(notification)).dump(); + wire = std::make_shared(nlohmann::json(std::move(notification)).dump()); + } + if (message_sender_) { + co_await message_sender_(*wire); + } else { + co_await transport_.write_message(*wire); } - co_await transport_.write_message(wire); } /** @@ -227,6 +242,7 @@ class Context { std::shared_ptr> cancelled_; std::optional progress_token_; const std::atomic* log_level_ = nullptr; + MessageSender message_sender_; }; } // namespace mcp diff --git a/include/mcp/detail/secure_random.hpp b/include/mcp/detail/secure_random.hpp new file mode 100644 index 0000000..d4c9974 --- /dev/null +++ b/include/mcp/detail/secure_random.hpp @@ -0,0 +1,9 @@ +#pragma once + +#include + +namespace mcp::detail { + +std::string generate_secure_session_id(); + +} // namespace mcp::detail diff --git a/include/mcp/detail/serialized_transport_writer.hpp b/include/mcp/detail/serialized_transport_writer.hpp new file mode 100644 index 0000000..2c1bd1e --- /dev/null +++ b/include/mcp/detail/serialized_transport_writer.hpp @@ -0,0 +1,63 @@ +#pragma once + +#include +#include + +#include + +#include +#include +#include +#include + +namespace mcp { + +class ITransport; + +namespace detail { + +struct SerializedTransportWriterState; + +/** + * @brief Serializes asynchronous writes to a shared transport. + * + * Calls are queued on an Asio strand and written in FIFO order. A failed + * transport write fails the active call, every queued call, and future calls. + * The shared implementation keeps both the transport and outstanding writes + * alive if this lightweight wrapper is destroyed. + */ +class MCP_API SerializedTransportWriter { + public: + SerializedTransportWriter(std::shared_ptr transport, + const boost::asio::any_io_executor& executor); + + /** + * @brief Queue one complete serialized message for writing. + * + * The input view is copied before this function returns its awaitable, so + * the caller does not need to keep the referenced storage alive. + */ + [[nodiscard]] Task write_message(std::string_view message) const; + + /** + * @brief Queue an already-owned message without another payload copy. + * @param message Immutable message storage retained until the write completes. + */ + [[nodiscard]] Task write_message(std::shared_ptr message) const; + + /** + * @brief Queue an owned message that may be canceled before its write starts. + * + * Once the transport write has started, setting @p cancel_before_start has + * no effect because ITransport has no per-write cancellation contract. + */ + [[nodiscard]] Task write_message( + std::shared_ptr message, + std::shared_ptr> cancel_before_start) const; + + private: + std::shared_ptr state_; +}; + +} // namespace detail +} // namespace mcp diff --git a/include/mcp/protocol/base.hpp b/include/mcp/protocol/base.hpp index bb12bec..e8432a6 100644 --- a/include/mcp/protocol/base.hpp +++ b/include/mcp/protocol/base.hpp @@ -4,12 +4,39 @@ #include #include #include +#include #include +#include #include #include namespace mcp { +namespace detail { + +inline void validate_jsonrpc_version(std::string_view version) { + if (version != "2.0") { + throw std::invalid_argument("jsonrpc must be \"2.0\""); + } +} + +/** + * @brief Reports whether @p key is present in @p json_obj and carries a value. + * + * An absent optional and one serialized as an explicit `null` are the same statement on + * the wire -- "no value" -- so `from_json` must read them the same way. Writing + * `contains(key)` alone accepts the null and then throws when the value is extracted. + * + * Returns false for a @p json_obj that is not an object, so a null payload decodes into a + * type whose members are all optional rather than failing. + */ +inline bool has_json_value(const nlohmann::json& json_obj, const char* key) { + const auto iter = json_obj.find(key); + return iter != json_obj.end() && !iter->is_null(); +} + +} // namespace detail + // MCP Protocol Constants /** @@ -22,6 +49,30 @@ inline constexpr int g_UNAUTHORIZED = -32000; */ inline constexpr int g_REQUEST_TIMEOUT = -32001; +/** + * @brief JSON-RPC-compatible client error code used when the transport closes. + */ +inline constexpr int g_CONNECTION_CLOSED = -32002; + +// Error-code allocation policy: -32000..-32019 is a grandfathered legacy band +// (pre-existing codes above); -32020..-32099 is the spec-reserved band for +// codes introduced by newer MCP spec revisions, such as the ones below. + +/** + * @brief JSON-RPC error code for a mismatched or missing required header. + */ +inline constexpr int g_HEADER_MISMATCH = -32020; + +/** + * @brief JSON-RPC error code for a missing required client capability. + */ +inline constexpr int g_MISSING_REQUIRED_CLIENT_CAPABILITY = -32021; + +/** + * @brief JSON-RPC error code for an unsupported protocol version. + */ +inline constexpr int g_UNSUPPORTED_PROTOCOL_VERSION = -32022; + /** * @brief JSON-RPC error code for invalid requests. */ @@ -67,9 +118,8 @@ NLOHMANN_JSON_SERIALIZE_ENUM(Role, {{Role::eUser, "user"}, {Role::eAssistant, "a /** * @brief JSON-RPC request identifier: either a string or an integer. * - * Wraps `std::variant` as a named type so that ADL-based - * `to_json`/`from_json` can live directly in `namespace mcp` without closing - * and reopening the namespace to specialize `nlohmann::adl_serializer`. + * Wraps `std::variant` as a named type so that `to_json`/`from_json` are found + * by ADL in `namespace mcp`. * * Implicit constructors allow transparent assignment from string and integer literals: * @code @@ -91,8 +141,7 @@ struct RequestId { bool operator==(const RequestId&) const = default; /** - * @brief Converts the request ID to a string representation (used as map key for pending request - * correlation). + * @brief Converts the request ID to its unadorned string representation. */ [[nodiscard]] std::string to_string() const { return std::visit( @@ -105,6 +154,19 @@ struct RequestId { }, value); } + + /** @brief Return a type-preserving key for request correlation. */ + [[nodiscard]] std::string correlation_key() const { + return std::visit( + [](const auto& val) -> std::string { + if constexpr (std::is_same_v, std::string>) { + return "s:" + val; + } else { + return "i:" + std::to_string(val); + } + }, + value); + } }; inline void to_json(nlohmann::json& json_obj, const RequestId& id) { @@ -142,8 +204,15 @@ inline void to_json(nlohmann::json& json_obj, const Error& error) { inline void from_json(const nlohmann::json& json_obj, Error& error) { json_obj.at("code").get_to(error.code); - json_obj.at("message").get_to(error.message); - if (json_obj.contains("data")) { + // JSON-RPC 2.0 requires `message`, but a peer that omits it or sends it as null has still + // told us which error occurred. Decoding to an empty message keeps `code` -- the part a + // caller acts on -- rather than discarding the whole error object. + if (detail::has_json_value(json_obj, "message")) { + json_obj.at("message").get_to(error.message); + } else { + error.message.clear(); + } + if (detail::has_json_value(json_obj, "data")) { error.data = json_obj.at("data"); } } @@ -181,7 +250,7 @@ inline void to_json(nlohmann::json& j, const RelatedTaskMetadata& t) { } inline void from_json(const nlohmann::json& j, RelatedTaskMetadata& t) { j.at("id").get_to(t.id); - if (j.contains("title")) { + if (detail::has_json_value(j, "title")) { t.title = j.at("title").get(); } } @@ -205,10 +274,10 @@ inline void to_json(nlohmann::json& json_obj, const TaskMetadata& meta) { } inline void from_json(const nlohmann::json& json_obj, TaskMetadata& meta) { - if (json_obj.contains("ttl")) { + if (detail::has_json_value(json_obj, "ttl")) { meta.ttl = json_obj.at("ttl").get(); } - if (json_obj.contains("relatedTasks")) { + if (detail::has_json_value(json_obj, "relatedTasks")) { meta.relatedTasks = json_obj.at("relatedTasks").get>(); } } @@ -236,8 +305,9 @@ inline void to_json(nlohmann::json& json_obj, const JSONRPCRequest& req) { inline void from_json(const nlohmann::json& json_obj, JSONRPCRequest& req) { req.id = json_obj.at("id").get(); json_obj.at("jsonrpc").get_to(req.jsonrpc); + detail::validate_jsonrpc_version(req.jsonrpc); json_obj.at("method").get_to(req.method); - if (json_obj.contains("params")) { + if (detail::has_json_value(json_obj, "params")) { req.params = json_obj.at("params"); } } @@ -260,8 +330,9 @@ inline void to_json(nlohmann::json& json_obj, const JSONRPCNotification& notif) inline void from_json(const nlohmann::json& json_obj, JSONRPCNotification& notif) { json_obj.at("jsonrpc").get_to(notif.jsonrpc); + detail::validate_jsonrpc_version(notif.jsonrpc); json_obj.at("method").get_to(notif.method); - if (json_obj.contains("params")) { + if (detail::has_json_value(json_obj, "params")) { notif.params = json_obj.at("params"); } } @@ -285,6 +356,7 @@ inline void to_json(nlohmann::json& json_obj, const JSONRPCResultResponse& resp) inline void from_json(const nlohmann::json& json_obj, JSONRPCResultResponse& resp) { resp.id = json_obj.at("id").get(); json_obj.at("jsonrpc").get_to(resp.jsonrpc); + detail::validate_jsonrpc_version(resp.jsonrpc); json_obj.at("result").get_to(resp.result); } @@ -294,12 +366,13 @@ inline void from_json(const nlohmann::json& json_obj, JSONRPCResultResponse& res struct JSONRPCErrorResponse { Error error; ///< The error object. std::string jsonrpc = "2.0"; ///< JSON-RPC version (always "2.0"). - std::optional id; ///< The request ID (may be absent for parse errors). + std::optional id; ///< The request ID, or no value when it cannot be determined. }; inline void to_json(nlohmann::json& json_obj, const JSONRPCErrorResponse& resp) { json_obj = nlohmann::json::object(); json_obj["error"] = resp.error; + json_obj["id"] = nullptr; json_obj["jsonrpc"] = resp.jsonrpc; if (resp.id) { json_obj["id"] = *resp.id; @@ -309,8 +382,11 @@ inline void to_json(nlohmann::json& json_obj, const JSONRPCErrorResponse& resp) inline void from_json(const nlohmann::json& json_obj, JSONRPCErrorResponse& resp) { json_obj.at("error").get_to(resp.error); json_obj.at("jsonrpc").get_to(resp.jsonrpc); - if (json_obj.contains("id")) { + detail::validate_jsonrpc_version(resp.jsonrpc); + if (detail::has_json_value(json_obj, "id")) { resp.id = json_obj.at("id").get(); + } else { + resp.id.reset(); } } @@ -322,7 +398,9 @@ inline void to_json(nlohmann::json& json_obj, const JSONRPCResponse& resp) { } inline void from_json(const nlohmann::json& json_obj, JSONRPCResponse& resp) { - if (json_obj.contains("error")) { + // A null `error` alongside a real `result` is how some peers spell "no error"; selecting on + // presence alone would route such a response into the error branch and fail to decode it. + if (detail::has_json_value(json_obj, "error")) { resp = json_obj.get(); } else { resp = json_obj.get(); @@ -338,7 +416,7 @@ inline void to_json(nlohmann::json& json_obj, const JSONRPCMessage& msg) { } inline void from_json(const nlohmann::json& json_obj, JSONRPCMessage& msg) { - if (json_obj.contains("error")) { + if (detail::has_json_value(json_obj, "error")) { msg = json_obj.get(); } else if (json_obj.contains("result")) { msg = json_obj.get(); diff --git a/include/mcp/protocol/capabilities.hpp b/include/mcp/protocol/capabilities.hpp index a6c47f4..de5fb1e 100644 --- a/include/mcp/protocol/capabilities.hpp +++ b/include/mcp/protocol/capabilities.hpp @@ -4,6 +4,7 @@ #include #include +#include #include #include #include @@ -32,6 +33,15 @@ constexpr std::string_view g_PROTOCOL_VERSION_2025_06_18 = "2025-06-18"; */ constexpr std::string_view g_PROTOCOL_VERSION_2025_11_25 = "2025-11-25"; +/** + * @brief Protocol version 2026-07-28 (stateless protocol revision). + * + * @details Deliberately absent from g_SUPPORTED_PROTOCOL_VERSIONS: legacy initialize + * negotiation must only resolve to versions whose semantics this SDK fully serves. This + * constant feeds the discovery surface until dual-era dispatch lands. + */ +constexpr std::string_view g_PROTOCOL_VERSION_2026_07_28 = "2026-07-28"; + /** * @brief The latest supported protocol version. */ @@ -44,6 +54,19 @@ constexpr std::array g_SUPPORTED_PROTOCOL_VERSIONS = { g_PROTOCOL_VERSION_2024_11_05, g_PROTOCOL_VERSION_2025_03_26, g_PROTOCOL_VERSION_2025_06_18, g_PROTOCOL_VERSION_2025_11_25}; +/** + * @brief Protocol versions advertised through the 2026-07-28 discovery surface. + * + * @details Consumed only by server/discover; deliberately distinct from + * g_SUPPORTED_PROTOCOL_VERSIONS, which alone governs legacy initialize negotiation. + * Advertising a version here makes no promise to the negotiation path: a legacy peer + * requesting 2026-07-28 still negotiates g_LATEST_PROTOCOL_VERSION until dual-era dispatch + * serves the new semantics. + */ +constexpr std::array g_DISCOVERABLE_PROTOCOL_VERSIONS = { + g_PROTOCOL_VERSION_2024_11_05, g_PROTOCOL_VERSION_2025_03_26, g_PROTOCOL_VERSION_2025_06_18, + g_PROTOCOL_VERSION_2025_11_25, g_PROTOCOL_VERSION_2026_07_28}; + /** * @brief Check whether a protocol version is supported by this SDK. * @@ -207,16 +230,16 @@ inline void to_json(nlohmann::json& json_obj, const Implementation& impl) { inline void from_json(const nlohmann::json& json_obj, Implementation& impl) { json_obj.at("name").get_to(impl.name); json_obj.at("version").get_to(impl.version); - if (json_obj.contains("title")) { + if (detail::has_json_value(json_obj, "title")) { impl.title = json_obj.at("title").get(); } - if (json_obj.contains("description")) { + if (detail::has_json_value(json_obj, "description")) { impl.description = json_obj.at("description").get(); } - if (json_obj.contains("websiteUrl")) { + if (detail::has_json_value(json_obj, "websiteUrl")) { impl.websiteUrl = json_obj.at("websiteUrl").get(); } - if (json_obj.contains("icons")) { + if (detail::has_json_value(json_obj, "icons")) { impl.icons = json_obj.at("icons").get>(); } } @@ -298,6 +321,8 @@ struct ClientCapabilities { std::optional roots; ///< Support for roots/list requests. std::optional sampling; ///< Support for sampling/createMessage requests. std::optional tasks; ///< Support for task lifecycle endpoints. + std::optional> + extensions; ///< Extension capability negotiation, keyed by extension name. }; inline void to_json(nlohmann::json& j, const ClientCapabilities::ElicitationCapability& cap) { @@ -445,6 +470,9 @@ inline void to_json(nlohmann::json& j, const ClientCapabilities& cap) { if (cap.tasks) { j["tasks"] = *cap.tasks; } + if (cap.extensions) { + j["extensions"] = *cap.extensions; + } } inline void from_json(const nlohmann::json& j, ClientCapabilities& cap) { @@ -463,6 +491,9 @@ inline void from_json(const nlohmann::json& j, ClientCapabilities& cap) { if (j.contains("tasks")) { cap.tasks = j.at("tasks").get(); } + if (j.contains("extensions")) { + cap.extensions = j.at("extensions").get>(); + } } /** @@ -528,6 +559,8 @@ struct ServerCapabilities { std::optional resources; ///< Support for resource endpoints. std::optional tasks; ///< Support for task lifecycle endpoints. std::optional tools; ///< Support for tool endpoints. + std::optional> + extensions; ///< Extension capability negotiation, keyed by extension name. }; inline void to_json(nlohmann::json& j, const ServerCapabilities::PromptsCapability& cap) { @@ -652,6 +685,9 @@ inline void to_json(nlohmann::json& j, const ServerCapabilities& cap) { if (cap.tools) { j["tools"] = *cap.tools; } + if (cap.extensions) { + j["extensions"] = *cap.extensions; + } } inline void from_json(const nlohmann::json& j, ServerCapabilities& cap) { @@ -676,6 +712,9 @@ inline void from_json(const nlohmann::json& j, ServerCapabilities& cap) { if (j.contains("tools")) { cap.tools = j.at("tools").get(); } + if (j.contains("extensions")) { + cap.extensions = j.at("extensions").get>(); + } } /** @@ -737,4 +776,105 @@ inline void from_json(const nlohmann::json& json_obj, InitializeResult& res) { } } +/** + * @brief Represents a server/discover request from the client. + * + * @details Carries no body parameters beyond the standard `_meta` field. Per the + * 2026-07-28 spec, `_meta` may include `io.modelcontextprotocol/protocolVersion`, + * `io.modelcontextprotocol/clientInfo`, and `io.modelcontextprotocol/clientCapabilities`; + * this type accepts and preserves the raw `_meta` object without interpreting its + * contents. + */ +struct DiscoverRequest { + std::optional meta; ///< Reserved for protocol use; contents not interpreted yet. +}; + +inline void to_json(nlohmann::json& json_obj, const DiscoverRequest& req) { + json_obj = nlohmann::json::object(); + if (req.meta) { + json_obj["_meta"] = *req.meta; + } +} + +inline void from_json(const nlohmann::json& json_obj, DiscoverRequest& req) { + if (json_obj.contains("_meta")) { + req.meta = json_obj.at("_meta"); + } +} + +/** + * @brief Caching scope hint for a cacheable result, per server/utilities/caching. + */ +enum class CacheScope : std::uint8_t { + ePublic, ///< The response does not contain user-specific data and may be shared. + ePrivate, ///< The response is scoped to the caller's authorization context. +}; + +NLOHMANN_JSON_SERIALIZE_ENUM(CacheScope, + {{CacheScope::ePublic, "public"}, {CacheScope::ePrivate, "private"}}) + +/** + * @brief Represents the result of a server/discover request. + * + * @details `resultType` is always "complete" for this result. + */ +struct DiscoverResult { + std::string resultType = "complete"; ///< Always "complete" for this result. + std::vector supportedVersions; ///< Protocol versions the server supports. + ServerCapabilities capabilities; ///< The capabilities supported by the server. + ServerInfo serverInfo; ///< Server identity, carried under `_meta` on the wire. + std::optional instructions; ///< Optional instructions for the client. + std::optional ttlMs; ///< Optional caching TTL hint, in milliseconds. + std::optional cacheScope; ///< Optional caching scope hint. +}; + +/** + * @brief Serializes DiscoverResult to JSON. + * + * @param json_obj The JSON object to populate. + * @param res The DiscoverResult object to serialize. + */ +inline void to_json(nlohmann::json& json_obj, const DiscoverResult& res) { + json_obj = nlohmann::json{ + {"resultType", res.resultType}, + {"supportedVersions", res.supportedVersions}, + {"capabilities", res.capabilities}, + {"_meta", {{"io.modelcontextprotocol/serverInfo", res.serverInfo}}}, + }; + if (res.instructions) { + json_obj["instructions"] = *res.instructions; + } + if (res.ttlMs) { + json_obj["ttlMs"] = *res.ttlMs; + } + if (res.cacheScope) { + json_obj["cacheScope"] = *res.cacheScope; + } +} + +/** + * @brief Deserializes DiscoverResult from JSON. + * + * @param json_obj The JSON object to read from. + * @param res The DiscoverResult object to populate. + */ +inline void from_json(const nlohmann::json& json_obj, DiscoverResult& res) { + json_obj.at("resultType").get_to(res.resultType); + json_obj.at("supportedVersions").get_to(res.supportedVersions); + json_obj.at("capabilities").get_to(res.capabilities); + if (json_obj.contains("_meta") && + json_obj.at("_meta").contains("io.modelcontextprotocol/serverInfo")) { + json_obj.at("_meta").at("io.modelcontextprotocol/serverInfo").get_to(res.serverInfo); + } + if (json_obj.contains("instructions")) { + res.instructions = json_obj.at("instructions").get(); + } + if (json_obj.contains("ttlMs")) { + res.ttlMs = json_obj.at("ttlMs").get(); + } + if (json_obj.contains("cacheScope")) { + res.cacheScope = json_obj.at("cacheScope").get(); + } +} + } // namespace mcp diff --git a/include/mcp/protocol/completion.hpp b/include/mcp/protocol/completion.hpp index 69ed665..d7a3554 100644 --- a/include/mcp/protocol/completion.hpp +++ b/include/mcp/protocol/completion.hpp @@ -106,7 +106,7 @@ inline void to_json(nlohmann::json& json_obj, const CompleteContext& context) { * @param context The CompleteContext object to populate. */ inline void from_json(const nlohmann::json& json_obj, CompleteContext& context) { - if (json_obj.contains("arguments")) { + if (detail::has_json_value(json_obj, "arguments")) { context.arguments = json_obj.at("arguments").get>(); } } diff --git a/include/mcp/protocol/content.hpp b/include/mcp/protocol/content.hpp index e088c0e..4556549 100644 --- a/include/mcp/protocol/content.hpp +++ b/include/mcp/protocol/content.hpp @@ -5,6 +5,16 @@ namespace mcp { +namespace detail { + +inline void validate_content_type(std::string_view actual, std::string_view expected) { + if (actual != expected) { + throw std::invalid_argument("content type must be \"" + std::string(expected) + "\""); + } +} + +} // namespace detail + /** * @brief Represents text content. */ @@ -141,6 +151,7 @@ inline void to_json(nlohmann::json& json_obj, const ResourceLink& link) { inline void from_json(const nlohmann::json& json_obj, ResourceLink& link) { json_obj.at("type").get_to(link.type); + detail::validate_content_type(link.type, "resource_link"); json_obj.at("uri").get_to(link.uri); json_obj.at("name").get_to(link.name); if (json_obj.contains("description")) { @@ -188,6 +199,7 @@ inline void to_json(nlohmann::json& json_obj, const EmbeddedResource& res) { inline void from_json(const nlohmann::json& json_obj, EmbeddedResource& res) { json_obj.at("type").get_to(res.type); + detail::validate_content_type(res.type, "resource"); json_obj.at("resource").get_to(res.resource); if (json_obj.contains("_meta")) { res.meta = json_obj.at("_meta").get(); @@ -218,6 +230,7 @@ inline void to_json(nlohmann::json& json_obj, const ToolUseContent& content) { inline void from_json(const nlohmann::json& json_obj, ToolUseContent& content) { json_obj.at("type").get_to(content.type); + detail::validate_content_type(content.type, "tool_use"); json_obj.at("id").get_to(content.id); json_obj.at("name").get_to(content.name); json_obj.at("input").get_to(content.input); @@ -254,15 +267,16 @@ inline void to_json(nlohmann::json& json_obj, const ToolResultContent& content) inline void from_json(const nlohmann::json& json_obj, ToolResultContent& content) { json_obj.at("type").get_to(content.type); + detail::validate_content_type(content.type, "tool_result"); json_obj.at("toolUseId").get_to(content.toolUseId); json_obj.at("content").get_to(content.content); - if (json_obj.contains("isError")) { + if (detail::has_json_value(json_obj, "isError")) { content.isError = json_obj.at("isError").get(); } - if (json_obj.contains("structuredContent")) { + if (detail::has_json_value(json_obj, "structuredContent")) { content.structuredContent = json_obj.at("structuredContent").get(); } - if (json_obj.contains("_meta")) { + if (detail::has_json_value(json_obj, "_meta")) { content.meta = json_obj.at("_meta").get(); } } @@ -285,6 +299,7 @@ inline void to_json(nlohmann::json& json_obj, const TextContent& content) { */ inline void from_json(const nlohmann::json& json_obj, TextContent& content) { json_obj.at("type").get_to(content.type); + detail::validate_content_type(content.type, "text"); json_obj.at("text").get_to(content.text); } @@ -307,6 +322,7 @@ inline void to_json(nlohmann::json& json_obj, const ImageContent& content) { */ inline void from_json(const nlohmann::json& json_obj, ImageContent& content) { json_obj.at("type").get_to(content.type); + detail::validate_content_type(content.type, "image"); json_obj.at("data").get_to(content.data); json_obj.at("mimeType").get_to(content.mimeType); } @@ -330,6 +346,7 @@ inline void to_json(nlohmann::json& json_obj, const AudioContent& content) { */ inline void from_json(const nlohmann::json& json_obj, AudioContent& content) { json_obj.at("type").get_to(content.type); + detail::validate_content_type(content.type, "audio"); json_obj.at("data").get_to(content.data); json_obj.at("mimeType").get_to(content.mimeType); } diff --git a/include/mcp/protocol/notification.hpp b/include/mcp/protocol/notification.hpp index d574b19..63c9bc8 100644 --- a/include/mcp/protocol/notification.hpp +++ b/include/mcp/protocol/notification.hpp @@ -42,7 +42,7 @@ inline void to_json(nlohmann::json& json_obj, const SetLevelRequestParams& param inline void from_json(const nlohmann::json& json_obj, SetLevelRequestParams& params) { json_obj.at("level").get_to(params.level); - if (json_obj.contains("_meta")) { + if (detail::has_json_value(json_obj, "_meta")) { params.meta = json_obj.at("_meta").get(); } } @@ -83,7 +83,7 @@ inline void to_json(nlohmann::json& j, const CancelledNotificationParams& params inline void from_json(const nlohmann::json& j, CancelledNotificationParams& params) { j.at("requestId").get_to(params.requestId); - if (j.contains("reason")) { + if (detail::has_json_value(j, "reason")) { params.reason = j.at("reason").get(); } } @@ -120,10 +120,10 @@ inline void to_json(nlohmann::json& j, const ProgressNotificationParams& params) inline void from_json(const nlohmann::json& j, ProgressNotificationParams& params) { j.at("progressToken").get_to(params.progressToken); j.at("progress").get_to(params.progress); - if (j.contains("total")) { + if (detail::has_json_value(j, "total")) { params.total = j.at("total").get(); } - if (j.contains("message")) { + if (detail::has_json_value(j, "message")) { params.message = j.at("message").get(); } } @@ -156,7 +156,7 @@ inline void to_json(nlohmann::json& j, const LoggingMessageNotificationParams& p inline void from_json(const nlohmann::json& j, LoggingMessageNotificationParams& params) { j.at("level").get_to(params.level); j.at("data").get_to(params.data); - if (j.contains("logger")) { + if (detail::has_json_value(j, "logger")) { params.logger = j.at("logger").get(); } } @@ -240,10 +240,10 @@ inline void to_json(nlohmann::json& j, const TaskStatusNotificationParams& p) { inline void from_json(const nlohmann::json& j, TaskStatusNotificationParams& p) { j.at("id").get_to(p.id); j.at("status").get_to(p.status); - if (j.contains("metadata")) { + if (detail::has_json_value(j, "metadata")) { p.metadata = j.at("metadata").get(); } - if (j.contains("message")) { + if (detail::has_json_value(j, "message")) { p.message = j.at("message").get(); } } diff --git a/include/mcp/protocol/prompts.hpp b/include/mcp/protocol/prompts.hpp index b2cccdc..283e68b 100644 --- a/include/mcp/protocol/prompts.hpp +++ b/include/mcp/protocol/prompts.hpp @@ -124,7 +124,7 @@ inline void to_json(nlohmann::json& json_obj, const GetPromptRequestParams& para inline void from_json(const nlohmann::json& json_obj, GetPromptRequestParams& params) { json_obj.at("name").get_to(params.name); - if (json_obj.contains("arguments")) { + if (detail::has_json_value(json_obj, "arguments")) { params.arguments = json_obj.at("arguments").get>(); } if (json_obj.contains("_meta")) { diff --git a/include/mcp/protocol/resources.hpp b/include/mcp/protocol/resources.hpp index a14525d..14df8f9 100644 --- a/include/mcp/protocol/resources.hpp +++ b/include/mcp/protocol/resources.hpp @@ -62,25 +62,25 @@ inline void to_json(nlohmann::json& json_obj, const Resource& resource) { inline void from_json(const nlohmann::json& json_obj, Resource& resource) { json_obj.at("uri").get_to(resource.uri); json_obj.at("name").get_to(resource.name); - if (json_obj.contains("description")) { + if (detail::has_json_value(json_obj, "description")) { resource.description = json_obj.at("description").get(); } - if (json_obj.contains("mimeType")) { + if (detail::has_json_value(json_obj, "mimeType")) { resource.mimeType = json_obj.at("mimeType").get(); } - if (json_obj.contains("_meta")) { + if (detail::has_json_value(json_obj, "_meta")) { resource.meta = json_obj.at("_meta").get(); } - if (json_obj.contains("annotations")) { + if (detail::has_json_value(json_obj, "annotations")) { resource.annotations = json_obj.at("annotations").get(); } - if (json_obj.contains("size")) { + if (detail::has_json_value(json_obj, "size")) { resource.size = json_obj.at("size").get(); } - if (json_obj.contains("title")) { + if (detail::has_json_value(json_obj, "title")) { resource.title = json_obj.at("title").get(); } - if (json_obj.contains("icons")) { + if (detail::has_json_value(json_obj, "icons")) { resource.icons = json_obj.at("icons").get>(); } } @@ -124,22 +124,22 @@ inline void to_json(nlohmann::json& json_obj, const ResourceTemplate& tmpl) { inline void from_json(const nlohmann::json& json_obj, ResourceTemplate& tmpl) { json_obj.at("uriTemplate").get_to(tmpl.uriTemplate); json_obj.at("name").get_to(tmpl.name); - if (json_obj.contains("description")) { + if (detail::has_json_value(json_obj, "description")) { tmpl.description = json_obj.at("description").get(); } - if (json_obj.contains("mimeType")) { + if (detail::has_json_value(json_obj, "mimeType")) { tmpl.mimeType = json_obj.at("mimeType").get(); } - if (json_obj.contains("_meta")) { + if (detail::has_json_value(json_obj, "_meta")) { tmpl.meta = json_obj.at("_meta").get(); } - if (json_obj.contains("annotations")) { + if (detail::has_json_value(json_obj, "annotations")) { tmpl.annotations = json_obj.at("annotations").get(); } - if (json_obj.contains("title")) { + if (detail::has_json_value(json_obj, "title")) { tmpl.title = json_obj.at("title").get(); } - if (json_obj.contains("icons")) { + if (detail::has_json_value(json_obj, "icons")) { tmpl.icons = json_obj.at("icons").get>(); } } diff --git a/include/mcp/protocol/tools.hpp b/include/mcp/protocol/tools.hpp index 206a168..e1581a6 100644 --- a/include/mcp/protocol/tools.hpp +++ b/include/mcp/protocol/tools.hpp @@ -1,5 +1,6 @@ #pragma once +#include #include #include #include @@ -131,25 +132,25 @@ inline void to_json(nlohmann::json& json_obj, const Tool& tool) { inline void from_json(const nlohmann::json& json_obj, Tool& tool) { json_obj.at("name").get_to(tool.name); json_obj.at("inputSchema").get_to(tool.inputSchema); - if (json_obj.contains("description")) { + if (detail::has_json_value(json_obj, "description")) { tool.description = json_obj.at("description").get(); } - if (json_obj.contains("_meta")) { + if (detail::has_json_value(json_obj, "_meta")) { tool.meta = json_obj.at("_meta").get(); } - if (json_obj.contains("annotations")) { + if (detail::has_json_value(json_obj, "annotations")) { tool.annotations = json_obj.at("annotations").get(); } - if (json_obj.contains("outputSchema")) { + if (detail::has_json_value(json_obj, "outputSchema")) { tool.outputSchema = json_obj.at("outputSchema").get(); } - if (json_obj.contains("title")) { + if (detail::has_json_value(json_obj, "title")) { tool.title = json_obj.at("title").get(); } - if (json_obj.contains("icons")) { + if (detail::has_json_value(json_obj, "icons")) { tool.icons = json_obj.at("icons").get>(); } - if (json_obj.contains("execution")) { + if (detail::has_json_value(json_obj, "execution")) { tool.execution = json_obj.at("execution").get(); } } @@ -158,9 +159,9 @@ inline void from_json(const nlohmann::json& json_obj, Tool& tool) { * @brief Parameters for calling a tool. */ struct CallToolParams { - std::string name; ///< The name of the tool to call. - nlohmann::json arguments; ///< The arguments for the tool call. - std::optional meta; ///< Reserved for protocol use. + std::string name; ///< The name of the tool to call. + nlohmann::json arguments = nlohmann::json::object(); ///< Optional tool arguments. + std::optional meta; ///< Reserved for protocol use. }; inline void to_json(nlohmann::json& json_obj, const CallToolParams& params) { @@ -172,8 +173,12 @@ inline void to_json(nlohmann::json& json_obj, const CallToolParams& params) { inline void from_json(const nlohmann::json& json_obj, CallToolParams& params) { json_obj.at("name").get_to(params.name); - json_obj.at("arguments").get_to(params.arguments); - if (json_obj.contains("_meta")) { + if (json_obj.contains("arguments") && !json_obj.at("arguments").is_object()) { + throw std::invalid_argument("CallToolParams arguments must be a JSON object"); + } + params.arguments = + json_obj.contains("arguments") ? json_obj.at("arguments") : nlohmann::json::object(); + if (detail::has_json_value(json_obj, "_meta")) { params.meta = json_obj.at("_meta").get(); } } @@ -188,6 +193,37 @@ struct CallToolResult { std::optional structuredContent; ///< Optional structured content. }; +/** + * @brief Create a successful tool result containing a text block. + * + * @param text Human-readable tool output. + * @return A protocol-valid CallToolResult. + */ +MCP_API CallToolResult make_tool_text_result(std::string text); + +/** + * @brief Create a failed tool result containing a text block. + * + * @param message Human-readable error description. + * @return A protocol-valid CallToolResult with isError set to true. + */ +MCP_API CallToolResult make_tool_error_result(std::string message); + +/** + * @brief Create a successful tool result containing structured output. + * + * structuredContent accepts any JSON value (object, array, string, number, + * boolean, or null), per MCP SEP-2106. A JSON serialization is also emitted + * as text for clients that do not consume structuredContent. + * + * @param structured_content Structured tool output. + * @param text Optional human-readable representation. When omitted, the + * structured value is serialized as JSON. + * @return A protocol-valid CallToolResult. + */ +MCP_API CallToolResult make_tool_structured_result(nlohmann::json structured_content, + std::optional text = std::nullopt); + /** * @brief Serializes CallToolResult to JSON. * @@ -215,13 +251,13 @@ inline void to_json(nlohmann::json& json_obj, const CallToolResult& result) { */ inline void from_json(const nlohmann::json& json_obj, CallToolResult& result) { json_obj.at("content").get_to(result.content); - if (json_obj.contains("isError")) { + if (detail::has_json_value(json_obj, "isError")) { result.isError = json_obj.at("isError").get(); } - if (json_obj.contains("_meta")) { + if (detail::has_json_value(json_obj, "_meta")) { result.meta = json_obj.at("_meta").get(); } - if (json_obj.contains("structuredContent")) { + if (detail::has_json_value(json_obj, "structuredContent")) { result.structuredContent = json_obj.at("structuredContent").get(); } } diff --git a/include/mcp/server/server.hpp b/include/mcp/server/server.hpp index ef3642d..ec946c7 100644 --- a/include/mcp/server/server.hpp +++ b/include/mcp/server/server.hpp @@ -44,6 +44,11 @@ using CompletionHandler = std::function(const CompleteParam namespace detail { +enum class ToolResultMode : std::uint8_t { + eNormalize, + eValidated, +}; + /// @brief Detects whether Fn is invocable with (In) returning Task. template concept AsyncHandlerNoCtx = requires(Fn fn, In in) { @@ -169,7 +174,10 @@ class MCP_API Server { tool.description = description; tool.inputSchema = input_schema; auto async_handler = detail::ensure_async_handler(std::move(handler)); - register_tool(tool, name, detail::wrap_handler(std::move(async_handler))); + constexpr auto result_mode = std::same_as + ? detail::ToolResultMode::eValidated + : detail::ToolResultMode::eNormalize; + register_tool(tool, name, detail::wrap_handler(std::move(async_handler)), result_mode); } /** @@ -187,7 +195,10 @@ class MCP_API Server { tool.inputSchema = input_schema; tool.outputSchema = output_schema; auto async_handler = detail::ensure_async_handler(std::move(handler)); - register_tool(tool, name, detail::wrap_handler(std::move(async_handler))); + constexpr auto result_mode = std::same_as + ? detail::ToolResultMode::eValidated + : detail::ToolResultMode::eNormalize; + register_tool(tool, name, detail::wrap_handler(std::move(async_handler)), result_mode); } /** @@ -206,6 +217,26 @@ class MCP_API Server { const nlohmann::json& input_schema, std::function handler); + /** + * @brief Register a low-level tool whose handler returns a complete CallToolResult JSON object. + * + * Unlike add_tool(), this escape hatch never normalizes arbitrary output. + * Every returned value is validated as a complete CallToolResult before it + * is sent to the client. + * + * @param tool Tool metadata. + * @param handler Asynchronous raw handler receiving tool arguments. + */ + void add_raw_tool(const Tool& tool, TypeErasedHandler handler); + + /** + * @brief Register a synchronous low-level tool returning complete CallToolResult JSON. + * + * @param tool Tool metadata. + * @param handler Synchronous raw handler receiving tool arguments. + */ + void add_raw_tool(const Tool& tool, std::function handler); + /** * @brief Register a resource with the server. * @@ -228,6 +259,26 @@ class MCP_API Server { */ void add_resource_template(const ResourceTemplate& tmpl); + /** + * @brief Register a resource template and its resources/read handler. + * + * Exact resources registered with add_resource() take precedence over + * template matches. Duplicate and structurally identical templates are + * rejected, and a URI matching more than one template is reported as an + * ambiguous request. + * + * @tparam In Handler input type, normally ReadResourceRequestParams. + * @tparam Out Handler result type, normally ReadResourceResult. + * @tparam Fn Handler callable (sync/async, with/without Context). + * @param tmpl Resource template metadata. + * @param handler Handler invoked when resources/read matches the template. + */ + template + void add_resource_template(const ResourceTemplate& tmpl, Fn handler) { + auto async_handler = detail::ensure_async_handler(std::move(handler)); + register_resource_template(tmpl, detail::wrap_handler(std::move(async_handler))); + } + /** * @brief Register a prompt with the server. * @@ -268,6 +319,36 @@ class MCP_API Server { */ void set_page_size(std::size_t size); + /** + * @brief Set the instructions string returned by initialize and server/discover. + * + * @details Read unsynchronized by request handlers; must be called before run(). + * + * @param instructions Natural-language guidance for LLMs on how to use this server. + */ + void set_instructions(std::string instructions); + + /** + * @brief Set the `ttlMs` caching hint returned by server/discover. + * + * @details Defaults to `0` ("do not cache") until overridden. Read unsynchronized by + * request handlers; must be called before run(). + * + * @param ttl_ms How long, in milliseconds, the client MAY consider the discover result fresh. + * @throws std::invalid_argument If `ttl_ms` is negative. + */ + void set_discover_ttl_ms(std::int64_t ttl_ms); + + /** + * @brief Set the `cacheScope` caching hint returned by server/discover. + * + * @details Defaults to `CacheScope::ePrivate` until overridden. Read unsynchronized by + * request handlers; must be called before run(). + * + * @param scope The intended cache scope of the discover result. + */ + void set_discover_cache_scope(CacheScope scope); + /** * @brief Send a JSON-RPC request to the connected client and await its response. * @@ -296,7 +377,9 @@ class MCP_API Server { /** * @brief Start the server session loop on the given transport. * - * Runs until the transport is closed or an error occurs. + * Runs until the transport is closed or an error occurs, then waits for + * request handlers already dispatched by this session to finish. The + * Server object must outlive the returned task. * * @param transport The transport to use for message exchange. Ownership is shared. * @param executor The executor to use for async operations. @@ -345,7 +428,11 @@ class MCP_API Server { * @brief Dispatch a parsed JSON-RPC request and return its serialized response. * * Used by stateless HTTP JSON mode to avoid creating a transport-backed - * server session for request/response messages. + * server session for request/response messages. The request envelope is + * validated, but no stateful initialize/initialized handshake is required. + * + * @details Request parameter decoding failures are returned as `g_INVALID_PARAMS` (-32602). + * Exceptions thrown by handlers or middleware are returned as `g_INTERNAL_ERROR` (-32603). * * @param json_msg The parsed JSON-RPC request object. * @return A task that resolves to the serialized JSON-RPC response. @@ -413,83 +500,13 @@ class MCP_API Server { [[nodiscard]] LoggingLevel get_log_level() const; private: - // Non-template registration helpers called by the template add_* methods above. - void register_tool(const Tool& tool, const std::string& name, TypeErasedHandler handler); + /// Non-template registration helpers called by the template add_* methods above. + void register_tool(const Tool& tool, const std::string& name, TypeErasedHandler handler, + detail::ToolResultMode result_mode); void register_resource(const Resource& resource, TypeErasedHandler handler); + void register_resource_template(const ResourceTemplate& tmpl, TypeErasedHandler handler); void register_prompt(const Prompt& prompt, TypeErasedHandler handler); - Context make_context(std::shared_ptr> cancelled = nullptr, - std::optional progress_token = std::nullopt); - - /** - * @brief Dispatches an incoming JSON-RPC request to the appropriate registered handler. - * - * @details Exceptions thrown during dispatch are caught and returned as `g_INTERNAL_ERROR` - * (-32603) JSON-RPC error responses. The exception's `what()` message becomes the error data. - * - * @param json_msg The raw JSON-RPC request object. - */ - Task dispatch_request(nlohmann::json json_msg); - - Task dispatch_request_wire(nlohmann::json json_msg); - - void dispatch_notification(const nlohmann::json& json_msg); - - void dispatch_response(const nlohmann::json& json_msg); - - Task handle_initialize(const nlohmann::json& json_msg); - Task handle_initialize_wire(const nlohmann::json& json_msg); - - Task handle_shutdown(const nlohmann::json& json_msg); - Task handle_shutdown_wire(const nlohmann::json& json_msg); - - Task handle_ping(const nlohmann::json& json_msg); - Task handle_ping_wire(const nlohmann::json& json_msg); - - Task handle_tools_call(const nlohmann::json& json_msg); - Task handle_tools_call_wire(const nlohmann::json& json_msg); - - Task handle_tools_list(const nlohmann::json& json_msg); - Task handle_tools_list_wire(const nlohmann::json& json_msg); - - Task handle_resources_list(const nlohmann::json& json_msg); - Task handle_resources_list_wire(const nlohmann::json& json_msg); - - Task handle_resources_read(const nlohmann::json& json_msg); - Task handle_resources_read_wire(const nlohmann::json& json_msg); - - Task handle_resource_templates_list(const nlohmann::json& json_msg); - Task handle_resource_templates_list_wire(const nlohmann::json& json_msg); - - Task handle_subscribe(const nlohmann::json& json_msg); - Task handle_subscribe_wire(const nlohmann::json& json_msg); - - Task handle_unsubscribe(const nlohmann::json& json_msg); - Task handle_unsubscribe_wire(const nlohmann::json& json_msg); - - Task handle_prompts_list(const nlohmann::json& json_msg); - Task handle_prompts_list_wire(const nlohmann::json& json_msg); - - Task handle_prompts_get(const nlohmann::json& json_msg); - Task handle_prompts_get_wire(const nlohmann::json& json_msg); - - Task handle_set_level(const nlohmann::json& json_msg); - Task handle_set_level_wire(const nlohmann::json& json_msg); - - Task handle_complete(const nlohmann::json& json_msg); - Task handle_complete_wire(const nlohmann::json& json_msg); - - Task invoke_tool_impl(CallToolParams params, - std::shared_ptr> cancelled, - std::optional progress_token); - - // [gcc11-sso: wire-builders] DO NOT convert to Task. - static std::string make_result_wire(const RequestId& id, nlohmann::json result); - static std::string make_error_wire(const RequestId& id, int code, std::string message); - - Task send_notification(const std::string& method, - const std::optional& params); - TypeErasedHandler build_middleware_chain(TypeErasedHandler final_handler); struct PaginationSlice { @@ -498,17 +515,17 @@ class MCP_API Server { std::optional next_cursor; }; - std::optional paginate(std::size_t total, const nlohmann::json& json_msg); - - [[nodiscard]] bool has_tool_output_schema(const std::string& name) const; - - void reset_session(); - struct PendingRequest; struct Session; + /// Tears down the current session. Runs on whatever thread destroys the Server, so it is safe + /// to call from any thread. + void reset_session(); + struct Impl; - std::unique_ptr impl_; + /// Shared rather than owned outright: work already in flight when this Server is destroyed + /// keeps the implementation alive until it finishes, so it never runs against freed state. + std::shared_ptr impl_; }; } // namespace mcp diff --git a/include/mcp/transport/http_client.hpp b/include/mcp/transport/http_client.hpp index 8d846e1..1a8bf05 100644 --- a/include/mcp/transport/http_client.hpp +++ b/include/mcp/transport/http_client.hpp @@ -5,11 +5,41 @@ #include #include +#include #include +#include #include +#include namespace mcp { +/** @brief HTTP-layer error returned by a remote MCP endpoint. */ +class MCP_API HttpStatusError : public std::runtime_error { + public: + /** + * @brief Construct an error for a non-success HTTP response. + * @param status Numeric HTTP response status. + * @param message Human-readable error description. + * @param authenticate_challenge Optional WWW-Authenticate header value. + */ + HttpStatusError(unsigned int status, std::string message, std::string authenticate_challenge = {}) + : std::runtime_error(std::move(message)), + status_(status), + authenticate_challenge_(std::move(authenticate_challenge)) {} + + /** @return Numeric HTTP response status. */ + [[nodiscard]] unsigned int status() const noexcept { return status_; } + + /** @return WWW-Authenticate header value, or an empty string when absent. */ + [[nodiscard]] const std::string& authenticate_challenge() const noexcept { + return authenticate_challenge_; + } + + private: + unsigned int status_; + std::string authenticate_challenge_; +}; + /** * @brief HTTP transport implementation for MCP client message exchange. * @@ -42,7 +72,7 @@ class MCP_API HttpClientTransport final : public ITransport { * @return The session identifier captured from server responses, or an empty string if no * session exists. */ - [[nodiscard]] const std::string& session_id() const; + [[nodiscard]] std::string session_id() const; /** * @brief Get the most recent SSE event identifier seen from the server. @@ -50,7 +80,16 @@ class MCP_API HttpClientTransport final : public ITransport { * @return The last event ID used for resumable HTTP replay, or an empty string if none was * received. */ - [[nodiscard]] const std::string& last_event_id() const; + [[nodiscard]] std::string last_event_id() const; + + /** + * @brief Configure a provider for HTTP Authorization: Bearer headers. + * + * Returning an empty string omits the header. Safe to call at any time, including while + * requests are in flight: each request pins the provider installed when it started, so a + * request already running keeps its provider and the next one picks up the new value. + */ + void set_bearer_token_provider(std::function provider); /** * @brief Dequeue the next MCP message received from the server. @@ -75,7 +114,7 @@ class MCP_API HttpClientTransport final : public ITransport { private: struct Impl; - std::unique_ptr impl_; + std::shared_ptr impl_; }; } // namespace mcp diff --git a/include/mcp/transport/http_server.hpp b/include/mcp/transport/http_server.hpp index 7e70ab9..2429391 100644 --- a/include/mcp/transport/http_server.hpp +++ b/include/mcp/transport/http_server.hpp @@ -14,6 +14,7 @@ #include #include #include +#include namespace mcp { @@ -179,11 +180,115 @@ class MCP_API HttpServerTransport final : public ITransport { * * When enabled, HTTP POST responses always use `application/json` and * outbound responses are not copied into the SSE replay event store. + * This option is atomic and may be changed while the listener is running. * * @param json_only True to bypass SSE framing and event storage. */ void set_json_only(bool json_only); + /** + * @brief Replace the allowlist used for requests carrying an Origin header. + * + * Requests without an Origin header remain valid for non-browser MCP clients. An Origin header + * is rejected by default until its exact value appears in this allowlist. + * Configure the allowlist before listen() or run() starts. + * + * @throws std::logic_error If listen() has already been called or the transport is closed. + */ + void set_allowed_origins(std::vector origins); + + /** + * @brief Explicitly opt into accepting every Origin header value. + * + * Configure this before listen() or run() starts. + * + * @throws std::logic_error If listen() has already been called or the transport is closed. + */ + void set_allow_all_origins(bool allow_all); + + /** + * @brief Require and validate an HTTP Authorization: Bearer header. + * + * Passing an empty validator disables HTTP authentication. + * Configure the validator before listen() or run() starts. + * + * @throws std::logic_error If listen() has already been called or the transport is closed. + */ + void set_bearer_token_validator(BearerTokenValidator validator); + + /** + * @brief Require an Authorization: Bearer header and validate it asynchronously. + * + * Use this when the decision needs I/O -- token introspection, a JWKS fetch -- so it suspends + * instead of blocking the executor that is concurrently serving MCP traffic. + * + * Passing an empty validator disables HTTP authentication. + * Configure the validator before listen() or run() starts. + * + * @throws std::logic_error If a synchronous validator is already installed, or if listen() has + * already been called or the transport is closed. + */ + void set_async_bearer_token_validator(AsyncBearerTokenValidator validator); + + /** + * @brief Cap the HTTP request body this transport will read. + * + * A request whose body exceeds the cap is answered `413 Payload Too Large` and its connection is + * closed; it never reaches MCP dispatch. Defaults to + * mcp::constants::g_default_max_request_body_bytes. Binary content travels as base64 inside the + * JSON body, so a cap near the size of the raw content rejects it. + * + * Configure the cap before listen() or run() starts. + * + * @throws std::invalid_argument If `max_bytes` is zero. + * @throws std::logic_error If listen() has already been called or the transport is closed. + */ + void set_max_request_body_bytes(std::size_t max_bytes); + + /** + * @brief Set the parameters sent in the `WWW-Authenticate` header of every 401. + * + * Without this call the transport sends the bare `Bearer` challenge. A client that has to + * discover where to obtain a token needs at least `resource_metadata`; setting protected + * resource metadata fills that field in automatically when it is left empty here. + * + * Configure the challenge before listen() or run() starts. + * + * @throws std::invalid_argument If a challenge value cannot be sent in a quoted-string. + * @throws std::logic_error If listen() has already been called or the transport is closed. + */ + void set_bearer_challenge(BearerChallengeConfig challenge); + + /** + * @brief Serve an RFC 9728 protected-resource metadata document. + * + * The document answers GET requests at its configured path without an Authorization header. When + * the bearer challenge carries no `resource_metadata`, it is populated with this document's URL. + * + * Configure the metadata before listen() or run() starts. + * + * @throws std::invalid_argument If `resource` is empty or is not an absolute URL. + * @throws std::logic_error If listen() has already been called or the transport is closed. + */ + void set_protected_resource_metadata(ProtectedResourceMetadataConfig metadata); + + /** + * @brief Exempt request paths from bearer validation and exclude them from MCP dispatch. + * + * An entry is excused from the bearer check and also removed from the set of paths MCP answers: + * if the protected-resource metadata route does not claim it, the request is answered `404 Not + * Found` rather than dispatched, so MCP is never served without authentication on an exempt path. + * + * Each entry is compared for equality against the path component of the request target, with any + * query string or fragment removed first, so `/health` also exempts `/health?probe=1`. Empty by + * default. + * + * Configure the paths before listen() or run() starts. + * + * @throws std::logic_error If listen() has already been called or the transport is closed. + */ + void set_unauthenticated_paths(std::vector paths); + /** * @brief Read the next queued JSON-RPC message from HTTP POST bodies. */ @@ -210,7 +315,8 @@ class MCP_API HttpServerTransport final : public ITransport { private: struct Impl; - std::unique_ptr impl_; + static Task listen_impl(std::shared_ptr impl); + std::shared_ptr impl_; }; } // namespace mcp diff --git a/include/mcp/transport/http_session_manager.hpp b/include/mcp/transport/http_session_manager.hpp index 40a25d4..690440e 100644 --- a/include/mcp/transport/http_session_manager.hpp +++ b/include/mcp/transport/http_session_manager.hpp @@ -15,6 +15,7 @@ #include #include #include +#include namespace mcp { @@ -40,8 +41,9 @@ namespace mcp { * auto server_factory = [](const asio::any_io_executor&) { * ServerCapabilities caps; * caps.tools = ServerCapabilities::ToolsCapability{}; - * Server server({"my-server", "1.0"}, std::move(caps)); - * server.add_tool("echo", "Echo tool", schema, handler); + * Implementation info{"my-server", "1.0"}; + * auto server = std::make_unique(info, std::move(caps)); + * server->add_tool("echo", "Echo tool", schema, handler); * return server; * }; * @@ -87,9 +89,112 @@ class MCP_API StreamableHttpSessionManager { * @param handler A function that receives the HTTP request and optionally returns * a response. If it returns std::nullopt, the request is handled * as MCP protocol. + * @throws std::logic_error If listen() has already been called or the manager is closed. */ void set_custom_request_handler(CustomRequestHandler handler); + /** + * @brief Replace the allowlist used for requests carrying an Origin header. + * + * Requests without an Origin header remain valid. Browser-originated requests are denied by + * default until their exact Origin value is present in this list. + * Configure the allowlist before listen() starts. + * + * @throws std::logic_error If listen() has already been called or the manager is closed. + */ + void set_allowed_origins(std::vector origins); + + /** + * @brief Explicitly opt into accepting every Origin header value before listen() starts. + * @throws std::logic_error If listen() has already been called or the manager is closed. + */ + void set_allow_all_origins(bool allow_all); + + /** + * @brief Require and validate an HTTP Authorization: Bearer header. + * + * Passing an empty validator disables HTTP authentication. + * Configure the validator before listen() starts. + * + * @throws std::logic_error If listen() has already been called or the manager is closed. + */ + void set_bearer_token_validator(BearerTokenValidator validator); + + /** + * @brief Require an Authorization: Bearer header and validate it asynchronously. + * + * Use this when the decision needs I/O -- token introspection, a JWKS fetch -- so it suspends + * instead of blocking the executor that is concurrently serving MCP traffic. + * + * Passing an empty validator disables HTTP authentication. + * Configure the validator before listen() starts. + * + * @throws std::logic_error If a synchronous validator is already installed, or if listen() has + * already been called or the manager is closed. + */ + void set_async_bearer_token_validator(AsyncBearerTokenValidator validator); + + /** + * @brief Cap the HTTP request body this manager will read. + * + * A request whose body exceeds the cap is answered `413 Payload Too Large` and its connection is + * closed; it never reaches MCP dispatch. Defaults to + * mcp::constants::g_default_max_request_body_bytes. Binary content travels as base64 inside the + * JSON body, so a cap near the size of the raw content rejects it. + * + * Configure the cap before listen() starts. + * + * @throws std::invalid_argument If `max_bytes` is zero. + * @throws std::logic_error If listen() has already been called or the manager is closed. + */ + void set_max_request_body_bytes(std::size_t max_bytes); + + /** + * @brief Set the parameters sent in the `WWW-Authenticate` header of every 401. + * + * Without this call the manager sends the bare `Bearer` challenge. A client that has to + * discover where to obtain a token needs at least `resource_metadata`; setting protected + * resource metadata fills that field in automatically when it is left empty here. + * + * Configure the challenge before listen() starts. + * + * @throws std::invalid_argument If a challenge value cannot be sent in a quoted-string. + * @throws std::logic_error If listen() has already been called or the manager is closed. + */ + void set_bearer_challenge(BearerChallengeConfig challenge); + + /** + * @brief Serve an RFC 9728 protected-resource metadata document. + * + * The document answers GET requests at its configured path without an Authorization header, ahead + * of both the bearer check and the custom request handler. When the bearer challenge carries no + * `resource_metadata`, it is populated with this document's URL. + * + * Configure the metadata before listen() starts. + * + * @throws std::invalid_argument If `resource` is empty or is not an absolute URL. + * @throws std::logic_error If listen() has already been called or the manager is closed. + */ + void set_protected_resource_metadata(ProtectedResourceMetadataConfig metadata); + + /** + * @brief Exempt request paths from bearer validation and exclude them from MCP dispatch. + * + * An entry is excused from the bearer check and also removed from the set of paths MCP answers: + * an exempt request still reaches the protected-resource metadata route and the custom request + * handler, but if both decline it is answered `404 Not Found` rather than dispatched, so MCP is + * never served without authentication on an exempt path. + * + * Each entry is compared for equality against the path component of the request target, with any + * query string or fragment removed first, so `/health` also exempts `/health?probe=1`. Empty by + * default. + * + * Configure the paths before listen() starts. + * + * @throws std::logic_error If listen() has already been called or the manager is closed. + */ + void set_unauthenticated_paths(std::vector paths); + /** * @brief Get the number of active sessions. * @@ -102,6 +207,7 @@ class MCP_API StreamableHttpSessionManager { * * When enabled, POST responses always use `application/json` and session * replay events are not stored. + * This option is atomic and may be changed while the listener is running. * * @param json_only True to bypass SSE framing and replay storage. */ @@ -119,6 +225,7 @@ class MCP_API StreamableHttpSessionManager { * behavior unchanged. * * @param enabled True to use stateless direct JSON handling. + * @throws std::logic_error If listen() has already been called or the manager is closed. */ void set_stateless_json_mode(bool enabled); @@ -130,6 +237,7 @@ class MCP_API StreamableHttpSessionManager { * If not set, tool handlers fall back to the HTTP executor (backward compatible). * * @param exec The executor to use for tool execution. + * @throws std::logic_error If listen() has already been called or the manager is closed. */ void set_tool_executor(const boost::asio::any_io_executor& exec); @@ -147,7 +255,8 @@ class MCP_API StreamableHttpSessionManager { private: struct Impl; - std::unique_ptr impl_; + static Task listen_impl(std::shared_ptr impl); + std::shared_ptr impl_; }; } // namespace mcp diff --git a/include/mcp/transport/http_types.hpp b/include/mcp/transport/http_types.hpp index e6b9c00..8e92562 100644 --- a/include/mcp/transport/http_types.hpp +++ b/include/mcp/transport/http_types.hpp @@ -1,12 +1,32 @@ #pragma once +#include +#include + #include +#include +#include #include +#include #include #include namespace mcp { +namespace constants { + +/** + * @brief Default cap on an HTTP request body accepted by the server transports, in bytes. + * + * Binary content reaches an MCP server as base64 inside the JSON body -- `ImageContent::data`, + * `BlobResourceContents::blob` and `AudioContent` have no out-of-band or streaming path -- and base64 + * inflates by 4/3. Eight mebibytes carries roughly 6 MB of raw content, while still bounding what one + * unauthenticated connection can make a server allocate. + */ +constexpr std::size_t g_default_max_request_body_bytes = 8 * 1024 * 1024; + +} // namespace constants + /** * @brief A generic list of string pairs, often used for query params or form fields. */ @@ -27,4 +47,112 @@ using StringRequest = boost::beast::http::request; +/// @brief Synchronous validator invoked at the HTTP boundary for a bearer token. +using BearerTokenValidator = std::function; + +/** + * @brief Asynchronous validator invoked at the HTTP boundary for a bearer token. + * + * Use this when deciding on a token requires I/O — token introspection, a JWKS fetch — so the + * decision suspends instead of blocking the executor that is concurrently serving MCP traffic. + * + * The token is passed by value: a view would point into the Beast request buffer, which is not + * guaranteed to outlive a suspension. + */ +using AsyncBearerTokenValidator = std::function(std::string)>; + +/** + * @brief Extract a bearer token from an HTTP Authorization header. + * @param authorization_header Complete Authorization header value. + * @return The characters after the Bearer scheme, or an empty view when the + * scheme is absent or no characters follow it. + */ +MCP_API std::string_view http_bearer_token(std::string_view authorization_header); + +/** + * @brief Extract the path component of an HTTP request target. + * @param target Request target, such as `/mcp?stream=1`. + * @return The characters before the first `?` or `#`, or the whole target when neither appears. + */ +MCP_API std::string_view http_request_path(std::string_view target); + +/** + * @brief Parameters advertised in the `WWW-Authenticate` challenge on a 401. + * + * Every field is optional. A default-constructed config renders the bare `Bearer` challenge, + * which is what the server transports send when no challenge is configured. + */ +struct BearerChallengeConfig { + std::string resource_metadata; ///< RFC 9728 metadata URL, sent as `resource_metadata="..."`. + std::string scope; ///< Space-delimited scopes, sent as `scope="..."`. + std::string realm; ///< RFC 7235 realm, sent as `realm="..."`. + std::string error; ///< RFC 6750 §3.1 code such as `invalid_token`. +}; + +/** + * @brief Render a challenge config as a `WWW-Authenticate` header value. + * + * Set parameters are emitted in the order realm, error, scope, resource_metadata, each as an + * RFC 7235 quoted-string with `\` and `"` backslash-escaped. + * + * @param config Challenge parameters to render. + * @return `Bearer` when no field is set, otherwise `Bearer key="value", ...`. + * @throws std::invalid_argument If a value holds a byte that cannot appear in a quoted-string, + * that is, anything outside horizontal tab and printable US-ASCII. The transports render + * the challenge once when it is configured, so an invalid value is reported to the caller + * that supplied it rather than while a request is being served. + */ +MCP_API std::string format_www_authenticate(const BearerChallengeConfig& config); + +/** + * @brief RFC 9728 protected-resource metadata served by a server transport. + */ +struct ProtectedResourceMetadataConfig { + std::string resource; ///< Required. Canonical resource URL. + std::vector authorization_servers; ///< Issuer URLs able to mint tokens. + std::vector scopes_supported; ///< Scopes the resource recognizes. + /// Path the document is served at. Empty means derive it from `resource` per RFC 9728 3.1; + /// set it only to override that derivation. + std::string path; +}; + +/** + * @brief Path the metadata document is served at. + * + * Returns `path` when it is set. Otherwise derives it the way RFC 9728 3.1 requires, by inserting the + * well-known segment between the authority and the resource's own path: a resource at + * `https://host/mcp` is described at `/.well-known/oauth-protected-resource/mcp`, and one at + * `https://host` at `/.well-known/oauth-protected-resource`. + * + * @param metadata Metadata whose `resource` supplies the path component. + * @return The path, always beginning with `/`. + * @throws std::invalid_argument If `resource` is not an absolute URL with an authority. + */ +MCP_API std::string protected_resource_metadata_path(const ProtectedResourceMetadataConfig& metadata); + +/** + * @brief Render protected-resource metadata as the JSON document clients fetch. + * + * Empty list members are omitted rather than written as empty arrays. + * + * @param metadata Metadata to render. + * @return A JSON object with `resource` and any populated list members. + */ +MCP_API std::string format_protected_resource_metadata(const ProtectedResourceMetadataConfig& metadata); + +/** + * @brief Build the absolute URL the metadata document is reachable at. + * + * The URL is the origin of `resource` — its scheme and authority — followed by + * protected_resource_metadata_path(). It is built entirely from `resource`, never from the + * address a transport happens to be bound to: behind a TLS terminator, a reverse proxy or a + * container port mapping the listener's own origin is not the one clients can reach, so inferring + * it would publish a URL that cannot be fetched. + * + * @param metadata Metadata whose `resource` supplies the origin and path. + * @return The absolute metadata URL. + * @throws std::invalid_argument If `resource` is not an absolute URL with an authority. + */ +MCP_API std::string protected_resource_metadata_url(const ProtectedResourceMetadataConfig& metadata); + } // namespace mcp diff --git a/include/mcp/transport/memory.hpp b/include/mcp/transport/memory.hpp index 4bab3c7..cff1711 100644 --- a/include/mcp/transport/memory.hpp +++ b/include/mcp/transport/memory.hpp @@ -1,25 +1,25 @@ #pragma once +#include #include -#include -#include -#include -#include +#include -#include - -#include #include -#include -#include #include #include #include -#include + +namespace boost::asio { +class any_io_executor; +} namespace mcp { +namespace detail { +struct MemoryTransportState; +} + /** * @brief In-memory transport implementation for testing. * @@ -27,172 +27,87 @@ namespace mcp { * written to one are available for reading on the other. This is useful * for testing MCP server/client interactions without real I/O. * - * Thread-safety is achieved via boost::asio::strand (no mutexes needed). - * The async pattern follows ScriptedTransport exactly: strand + timer + queue. + * Endpoint state is retained by pending operations, so destroying a + * MemoryTransport wrapper does not invalidate an in-flight read or write. + * Once the wrapper and its final in-flight operation are gone, the peer is + * closed and its pending readers are woken. + * Queue and pending-read state is serialized by an internal Boost.Asio strand. + * Concurrent reads use independent wake-up operations and do not cancel each other. * * @note Always create instances via create_memory_transport_pair(). Never - * construct MemoryTransport directly — the peer link would be unset. + * construct MemoryTransport directly unless set_peer() is called before use. */ -class MemoryTransport final : public ITransport { +class MCP_API MemoryTransport final : public ITransport { public: /** * @brief Constructs a MemoryTransport with the given executor. * * @param executor The executor for async operations. */ - explicit MemoryTransport(const boost::asio::any_io_executor& executor) - : strand_(boost::asio::make_strand(executor)), timer_(strand_) { - timer_.expires_at(std::chrono::steady_clock::time_point::max()); - } + explicit MemoryTransport(const boost::asio::any_io_executor& executor); + ~MemoryTransport() override; + + MemoryTransport(const MemoryTransport&) = delete; + MemoryTransport& operator=(const MemoryTransport&) = delete; + MemoryTransport(MemoryTransport&&) = delete; + MemoryTransport& operator=(MemoryTransport&&) = delete; /** * @brief Reads the next message from the internal queue. * - * Suspends using a timer until a message is available or the transport is closed. - * Follows the exact pattern from ScriptedTransport. + * Suspends until a message is available or the transport is closed. * * @return A task yielding the next queued message. * @throws std::runtime_error If the transport is closed. */ - Task read_message() override { - if (strand_.running_in_this_thread() && !closed_ && !incoming_.empty()) { - co_return pop_message_as_string(); - } - - co_await boost::asio::post(strand_, boost::asio::use_awaitable); - - for (;;) { - if (closed_) { - throw std::runtime_error("transport closed"); - } - if (!incoming_.empty()) { - co_return pop_message_as_string(); - } - timer_.expires_at(std::chrono::steady_clock::time_point::max()); - try { - co_await timer_.async_wait(boost::asio::use_awaitable); - } catch (const boost::system::system_error& err) { - if (err.code() != boost::asio::error::operation_aborted) { - throw; - } - } - } - } - - Task read_json() { - if (strand_.running_in_this_thread() && !closed_ && !incoming_.empty()) { - co_return pop_message_as_json(); - } - - co_await boost::asio::post(strand_, boost::asio::use_awaitable); - - for (;;) { - if (closed_) { - throw std::runtime_error("transport closed"); - } - if (!incoming_.empty()) { - co_return pop_message_as_json(); - } - timer_.expires_at(std::chrono::steady_clock::time_point::max()); - try { - co_await timer_.async_wait(boost::asio::use_awaitable); - } catch (const boost::system::system_error& err) { - if (err.code() != boost::asio::error::operation_aborted) { - throw; - } - } - } - } - - Task write_json(nlohmann::json message) { - co_await boost::asio::post(strand_, boost::asio::use_awaitable); - if (closed_) { - throw std::runtime_error("transport closed"); - } - auto peer = peer_.lock(); - if (!peer) { - throw std::runtime_error("peer transport not set"); - } - boost::asio::post(peer->strand_, [peer, msg = std::move(message)]() mutable { - peer->incoming_.emplace(std::move(msg)); - peer->timer_.cancel(); - }); - co_return; - } + Task read_message() override; + + /** + * @brief Reads and parses the next message as JSON. + * + * @return A task yielding the next queued JSON value. + * @throws std::runtime_error If the transport is closed. + */ + Task read_json(); + + /** + * @brief Writes a JSON value to the peer's queue. + * + * @param message The JSON value to send. + * @return A task that completes when the peer queue has been updated. + * @throws std::runtime_error If either endpoint is closed or the peer is unavailable. + */ + Task write_json(nlohmann::json message); /** * @brief Writes a message to the peer's queue and wakes the peer. * + * The message bytes are copied before this function returns, so a caller may + * safely modify or destroy the storage behind @p message before awaiting the task. + * * @param message The message to send. * @return A task that completes when the peer queue has been updated. - * @throws std::runtime_error If the transport is closed or the peer is not set or has been - * destroyed. + * @throws std::runtime_error If either endpoint is closed or the peer is unavailable. */ - Task write_message(std::string_view message) override { - co_await boost::asio::post(strand_, boost::asio::use_awaitable); - if (closed_) { - throw std::runtime_error("transport closed"); - } - auto peer = peer_.lock(); - if (!peer) { - throw std::runtime_error("peer transport not set"); - } - // Post message to peer's strand - boost::asio::post(peer->strand_, [peer, msg = std::string(message)]() mutable { - peer->incoming_.emplace(std::move(msg)); - peer->timer_.cancel(); - }); - } + Task write_message(std::string_view message) override; /** - * @brief Closes the transport and the peer. + * @brief Closes the transport and its peer. * - * Sets closed flag, cancels pending timers, and closes the peer transport. + * Safe to call concurrently and multiple times. Pending readers on both + * endpoints are woken and future reads and writes fail. */ - void close() override { - if (closed_) { - return; - } - closed_ = true; - timer_.cancel(); - if (auto peer = peer_.lock(); peer && !peer->closed_) { - peer->close(); - } - } + void close() override; /** * @brief Sets the peer transport for bidirectional communication. * - * @param peer Shared pointer to the peer MemoryTransport. + * @param peer Shared pointer to the peer MemoryTransport, or null to unlink it. */ - void set_peer(const std::shared_ptr& peer) { peer_ = peer; } + void set_peer(const std::shared_ptr& peer); private: - using MemoryMessage = std::variant; - - std::string pop_message_as_string() { - auto message = std::move(incoming_.front()); - incoming_.pop(); - if (auto* raw = std::get_if(&message)) { - return std::move(*raw); - } - return std::get(message).dump(); - } - - nlohmann::json pop_message_as_json() { - auto message = std::move(incoming_.front()); - incoming_.pop(); - if (auto* json = std::get_if(&message)) { - return std::move(*json); - } - return nlohmann::json::parse(std::get(message)); - } - - boost::asio::strand strand_; - boost::asio::steady_timer timer_; - std::queue incoming_; - std::weak_ptr peer_; - bool closed_ = false; + std::shared_ptr state_; }; /** @@ -201,19 +116,13 @@ class MemoryTransport final : public ITransport { * Messages written to the first transport are readable on the second, * and vice versa. Both transports share the same executor. * - * The peer relationship uses weak_ptr to avoid reference cycles; each - * transport holds only a weak reference to the other. + * The peer relationship uses weak ownership to avoid reference cycles. + * Pending operations retain only the endpoint states they need. * * @param executor The executor for both transports. * @return A pair of connected transport instances that can exchange in-memory messages. */ -inline std::pair, std::shared_ptr> create_memory_transport_pair( - const boost::asio::any_io_executor& executor) { - auto transport_a = std::make_shared(executor); - auto transport_b = std::make_shared(executor); - transport_a->set_peer(transport_b); - transport_b->set_peer(transport_a); - return {transport_a, transport_b}; -} +MCP_API std::pair, std::shared_ptr> +create_memory_transport_pair(const boost::asio::any_io_executor& executor); } // namespace mcp diff --git a/include/mcp/transport/stdio.hpp b/include/mcp/transport/stdio.hpp index 8a46cb2..35a6e4f 100644 --- a/include/mcp/transport/stdio.hpp +++ b/include/mcp/transport/stdio.hpp @@ -20,6 +20,12 @@ namespace mcp { * Reads JSON-RPC messages line-by-line from an input stream and writes * responses to an output stream. By default uses std::cin / std::cout, * but accepts arbitrary streams for testing. + * + * @warning On the stdio transport the output stream *is* the protocol channel. + * With the default std::cout, any other write to stdout (a stray printf, a + * logging library's default sink) lands between framed messages and the peer's + * parser rejects it; the transport cannot detect this. Applications that cannot + * guarantee a silent stdout should use create_owning_stdout(). */ class MCP_API StdioTransport final : public ITransport { public: @@ -29,10 +35,44 @@ class MCP_API StdioTransport final : public ITransport { * @param executor The executor to use for async operations. * @param input Input stream to read messages from (default: std::cin). * @param output Output stream to write messages to (default: std::cout). + * + * @note The input and output streams must outlive the transport. Once a + * read has started, destruction waits for the blocking input operation to + * finish. Because a generic std::istream cannot be cancelled, callers + * using a blocking custom stream must make it return EOF before destroying + * the transport. */ explicit StdioTransport(const boost::asio::any_io_executor& executor, std::istream& input = std::cin, std::ostream& output = std::cout); + /** + * @brief Construct a StdioTransport that owns the process's standard output. + * + * Duplicates the current standard output onto a private descriptor, writes + * the protocol there, and points the process's standard output at standard + * error, so anything the application writes to stdout (std::cout, printf, a + * logging library's default sink) lands on stderr instead of in the message + * stream. The original standard output is restored when the transport is + * destroyed. + * + * @param executor The executor to use for async operations. + * @param input Input stream to read messages from (default: std::cin). + * + * @return The transport. It must be destroyed before the process relies on + * stdout again. + * + * @throws std::runtime_error if standard output cannot be duplicated or + * redirected, or if another StdioTransport in this process already owns it. + * + * @note This changes process-global state, so construct the transport + * during start-up, before other threads write to stdout, and keep at most + * one owning transport alive at a time. On Windows the redirect covers the + * C runtime (std::cout, printf, fwrite); code writing directly to the + * Win32 STD_OUTPUT_HANDLE is not covered. + */ + static std::unique_ptr create_owning_stdout( + const boost::asio::any_io_executor& executor, std::istream& input = std::cin); + ~StdioTransport() override; StdioTransport(const StdioTransport&) = delete; @@ -43,6 +83,9 @@ class MCP_API StdioTransport final : public ITransport { /** * @brief Read the next newline-delimited message from the input stream. * + * At most one read may be outstanding at a time. A concurrent read + * completes with std::logic_error. + * * Throws std::runtime_error if the transport is closed or the input stream * reaches EOF. */ @@ -59,10 +102,16 @@ class MCP_API StdioTransport final : public ITransport { /** * @brief Close the transport. Safe to call multiple times. + * + * Wakes the asynchronous reader, but cannot interrupt a blocking operation + * inside the supplied std::istream. */ void close() override; private: + struct OwnStdoutTag {}; + StdioTransport(const boost::asio::any_io_executor& executor, std::istream& input, OwnStdoutTag); + struct Impl; std::unique_ptr impl_; }; diff --git a/include/mcp/transport/websocket.hpp b/include/mcp/transport/websocket.hpp index d97ded3..eee5b68 100644 --- a/include/mcp/transport/websocket.hpp +++ b/include/mcp/transport/websocket.hpp @@ -6,6 +6,7 @@ #include #include +#include #include #include #include @@ -38,13 +39,16 @@ class MCP_API WebSocketServerTransport final : public ITransport { /** * @brief Read the next message from the WebSocket connection. * - * Throws std::runtime_error if the transport is closed or the connection drops. + * @throws std::logic_error If another read is already outstanding. + * @throws std::runtime_error If the transport is closed or the connection drops. */ Task read_message() override; /** * @brief Send a message over the WebSocket connection. * + * Concurrent writes are serialized by the transport. + * * @param message The message to send. */ Task write_message(std::string_view message) override; @@ -56,7 +60,7 @@ class MCP_API WebSocketServerTransport final : public ITransport { private: struct Impl; - std::unique_ptr impl_; + std::shared_ptr impl_; }; /** @@ -74,9 +78,14 @@ class MCP_API WebSocketClientTransport final : public ITransport { * @param host Remote hostname or IP address. * @param port Remote port number. * @param path WebSocket request path (default: "/"). + * @param connect_timeout Limit on the whole connection setup, from the TCP connect to the + * end of the WebSocket handshake. A peer that accepts the connection and + * then stalls fails the pending call after this long instead of hanging it. + * Zero disables the limit. Established connections are not subject to it. */ WebSocketClientTransport(const boost::asio::any_io_executor& executor, std::string host, - std::string port, std::string path = "/"); + std::string port, std::string path = "/", + std::chrono::milliseconds connect_timeout = std::chrono::seconds(30)); ~WebSocketClientTransport() override; @@ -89,7 +98,9 @@ class MCP_API WebSocketClientTransport final : public ITransport { * @brief Read the next message from the WebSocket connection. * * Connects and performs the WebSocket handshake on the first call. - * Throws std::runtime_error if the transport is closed or the connection drops. + * + * @throws std::logic_error If another read is already outstanding. + * @throws std::runtime_error If the transport is closed or the connection drops. */ Task read_message() override; @@ -97,6 +108,7 @@ class MCP_API WebSocketClientTransport final : public ITransport { * @brief Send a message over the WebSocket connection. * * Connects and performs the WebSocket handshake on the first call. + * Concurrent writes are serialized by the transport. * * @param message The message to send. */ @@ -109,7 +121,7 @@ class MCP_API WebSocketClientTransport final : public ITransport { private: struct Impl; - std::unique_ptr impl_; + std::shared_ptr impl_; }; } // namespace mcp diff --git a/packaging/debian/rules.in b/packaging/debian/rules.in index 9757b27..f50436b 100644 --- a/packaging/debian/rules.in +++ b/packaging/debian/rules.in @@ -7,7 +7,7 @@ export DEB_CXXFLAGS_MAINT_APPEND = -fvisibility=hidden -fvisibility-inlines-hidd dh $@ --buildsystem=cmake+ninja override_dh_auto_configure: - dh_auto_configure -- -DMCP_CPP_SDK_BUILD_SHARED=ON -DMCP_CPP_SDK_BUILD_STATIC=ON -DBUILD_TESTING=ON -DBUILD_EXAMPLES=OFF -DBUILD_DOCS=OFF -DCMAKE_BUILD_TYPE=RelWithDebInfo + dh_auto_configure -- -DMCP_CPP_SDK_BUILD_SHARED=ON -DMCP_CPP_SDK_BUILD_STATIC=ON -DBUILD_TESTING=ON -DMCP_CPP_SDK_CHECK_JSON_MATRIX=OFF -DBUILD_EXAMPLES=OFF -DBUILD_DOCS=OFF -DCMAKE_BUILD_TYPE=RelWithDebInfo override_dh_auto_test: CTEST_OUTPUT_ON_FAILURE=1 dh_auto_test diff --git a/scripts/build.py b/scripts/build.py index 6070cdd..07377ed 100644 --- a/scripts/build.py +++ b/scripts/build.py @@ -85,7 +85,7 @@ def ensure_conan_profile(): print("[+] Conan profile created") -def conan_install(output_folder, jobs, cppstd, build_type="Release"): +def conan_install(output_folder, jobs, cppstd, build_type="Release", extra_args=()): ensure_conan_profile() run( "conan", "install", ".", @@ -95,6 +95,7 @@ def conan_install(output_folder, jobs, cppstd, build_type="Release"): "-s", f"build_type={build_type}", "-c", "tools.cmake.cmaketoolchain:generator=Ninja", "-c", f"tools.build:jobs={jobs}", + *extra_args, ) generators = Path(output_folder) / "build" / build_type / "generators" toolchain = generators / "conan_toolchain.cmake" @@ -129,8 +130,126 @@ def cmake_configure(build_dir, build_type, toolchain, *extra_args): ) -def cmake_build(build_dir, jobs): - run("cmake", "--build", build_dir, f"-j{jobs}") +def no_aslr(): + """Command prefix that disables ASLR, which ThreadSanitizer and MemorySanitizer need on kernels + with high mmap entropy. + + The test binary also runs at build time (gtest_discover_tests), so the build and ctest both use it. + """ + setarch = shutil.which("setarch") + if setarch and sys.platform.startswith("linux"): + return [setarch, os.uname().machine, "-R"] + return [] + + +def cmake_build(build_dir, jobs, prefix=()): + run(*prefix, "cmake", "--build", build_dir, f"-j{jobs}") + + +# The Clang release the Clang builds look for first, and the LLVM commit (llvmorg-18.1.8) whose +# libc++ the MemorySanitizer build compiles with it. Keep the two on the same major version. +CLANG_VERSION = "18" +MSAN_LLVM_COMMIT = "3b5b5c1ec4a3095ab096dd780e84d7ab81f3d7ff" +MSAN_LIBCXX_DIR = "build/msan-libcxx" + + +def find_clang(): + """Return (C compiler, C++ compiler) for Clang, or None if there is no complete pair.""" + for suffix in (f"-{CLANG_VERSION}", ""): + c_compiler = shutil.which(f"clang{suffix}") + cxx_compiler = shutil.which(f"clang++{suffix}") + if c_compiler and cxx_compiler: + return c_compiler, cxx_compiler + return None + + +def build_msan_libcxx(prefix, clang, jobs): + """Build libc++ and libc++abi instrumented for MemorySanitizer and install them under prefix. + + MemorySanitizer tracks which bytes have been initialised, and only instrumented code tells it. + Memory a prebuilt standard library writes stays "uninitialised" to it and every later read is + reported, so the standard library has to be an instrumented build as well. + """ + prefix = Path(prefix).resolve() + # Written last, so it also marks a build that ran to the end. A prefix from another commit, + # such as a stale CI cache, is rebuilt. + stamp = prefix / "llvm-commit" + if stamp.is_file() and stamp.read_text().strip() == MSAN_LLVM_COMMIT: + return prefix + + work = prefix.with_name(prefix.name + "-work") + shutil.rmtree(work, ignore_errors=True) + shutil.rmtree(prefix, ignore_errors=True) + source = work / "llvm-project" + source.mkdir(parents=True) + # Only the runtimes and the CMake modules they share with LLVM, at one pinned commit. + run("git", "init", "-q", str(source)) + run("git", "-C", str(source), "remote", "add", "origin", + "https://github.com/llvm/llvm-project.git") + run("git", "-C", str(source), "sparse-checkout", "set", + "runtimes", "libcxx", "libcxxabi", "llvm/cmake", "llvm/utils/llvm-lit", "cmake", + "third-party") + run("git", "-C", str(source), "fetch", "-q", "--depth", "1", "--filter=blob:none", + "origin", MSAN_LLVM_COMMIT) + run("git", "-C", str(source), "-c", "advice.detachedHead=false", "checkout", "-q", "FETCH_HEAD") + + c_compiler, cxx_compiler = clang + launcher_args = [] + launcher = compiler_launcher() + if launcher: + launcher_args = [f"-DCMAKE_C_COMPILER_LAUNCHER={launcher}", + f"-DCMAKE_CXX_COMPILER_LAUNCHER={launcher}"] + run( + "cmake", "-S", str(source / "runtimes"), "-B", str(work / "build"), "-G", "Ninja", + "-DCMAKE_BUILD_TYPE=Release", + f"-DCMAKE_C_COMPILER={c_compiler}", + f"-DCMAKE_CXX_COMPILER={cxx_compiler}", + f"-DCMAKE_INSTALL_PREFIX={prefix}", + "-DLLVM_ENABLE_RUNTIMES=libcxx;libcxxabi", + "-DLLVM_USE_SANITIZER=MemoryWithOrigins", + # Unwind with the system's libgcc. An instrumented libunwind reads registers that its own + # assembly saved, reports them as uninitialised, and unwinds again to print the report. + "-DLIBCXXABI_USE_LLVM_UNWINDER=OFF", + "-DLIBCXX_INCLUDE_TESTS=OFF", + "-DLIBCXX_INCLUDE_BENCHMARKS=OFF", + "-DLIBCXXABI_INCLUDE_TESTS=OFF", + *launcher_args, + ) + run("cmake", "--build", str(work / "build"), f"-j{jobs}", + "--target", "install-cxx", "install-cxxabi") + shutil.rmtree(work, ignore_errors=True) + stamp.write_text(MSAN_LLVM_COMMIT + "\n") + return prefix + + +def msan_conan_args(libcxx, clang): + """Conan arguments that build the SDK and every dependency against the instrumented libc++. + + The flags reach the dependencies through Conan and the SDK through the toolchain file it + generates, so CMake needs no option of its own for this mode. + """ + c_compiler, cxx_compiler = clang + compile_flags = ["-fsanitize=memory", "-fsanitize-memory-track-origins=2", + "-fno-omit-frame-pointer", "-g"] + # Conan adds -stdlib=libc++ for the libc++ setting. -nostdinc++ then makes the instrumented + # headers the only ones, which leaves -stdlib with nothing to do when compiling. + quiet = "-Wno-unused-command-line-argument" + cxx_flags = compile_flags + ["-nostdinc++", f"-isystem{libcxx}/include/c++/v1", quiet] + link_flags = ["-fsanitize=memory", "-stdlib=libc++", f"-L{libcxx}/lib", + f"-Wl,-rpath,{libcxx}/lib", "-lc++abi", quiet] + executables = {"c": c_compiler, "cpp": cxx_compiler} + return [ + "-s", "compiler=clang", + "-s", f"compiler.version={CLANG_VERSION}", + "-s", "compiler.libcxx=libc++", + # MemorySanitizer cannot see what OpenSSL's hand-written assembly initialises. + "-o", "openssl/*:no_asm=True", + "-c", f"tools.build:compiler_executables={json.dumps(executables)}", + "-c", f"tools.build:cflags={json.dumps(compile_flags)}", + "-c", f"tools.build:cxxflags={json.dumps(cxx_flags)}", + "-c", f"tools.build:exelinkflags={json.dumps(link_flags)}", + "-c", f"tools.build:sharedlinkflags={json.dumps(link_flags)}", + ] def write_user_presets(generators_dir): @@ -153,12 +272,16 @@ def write_user_presets(generators_dir): python scripts/build.py --test release build + run tests python scripts/build.py --debug --test debug build + run tests python scripts/build.py --sanitize --test ASan/UBSan build + run tests + python scripts/build.py --tsan --test ThreadSanitizer build + run tests + python scripts/build.py --tsan --compiler clang --test same, compiled with Clang + python scripts/build.py --msan --test MemorySanitizer build + run tests python scripts/build.py --coverage --test gcov build + run tests + report python scripts/build.py --test skip examples, run tests python scripts/build.py --sanitize --test sanitized tests, no examples python scripts/build.py --linkage shared --test build/test only shared SDK python scripts/build.py --linkage static --test build/test only static SDK python scripts/build.py --cppstd 23 --test build/test a C++23 consumer + python scripts/build.py --conformance build conformance fixtures python scripts/build.py --docs build documentation python scripts/build.py --clean remove all build artifacts """ @@ -174,12 +297,23 @@ def main(): help="Debug build (default: Release)") parser.add_argument("--sanitize", action="store_true", help="Enable ASan + UBSan (implies --debug)") + parser.add_argument("--tsan", action="store_true", + help="Enable ThreadSanitizer (implies --debug; excludes --sanitize)") + parser.add_argument("--msan", action="store_true", + help="Enable MemorySanitizer (implies --debug and Clang; excludes the " + "other sanitizers). Builds an instrumented libc++ on first use and " + "rebuilds the Conan dependencies against it") + parser.add_argument("--compiler", choices=("default", "clang"), default="default", + help="Compile the SDK and tests with Clang (Conan dependencies keep " + "the detected profile); the build directory gains a -clang suffix") parser.add_argument("--coverage", action="store_true", help="Enable gcov coverage (implies --debug)") parser.add_argument("--test", action="store_true", help="Build and run tests") parser.add_argument("--examples", action="store_true", help="Build example programs") + parser.add_argument("--conformance", action="store_true", + help="Build fixtures for the official MCP conformance runner") parser.add_argument( "--linkage", choices=("both", "shared", "static"), @@ -206,17 +340,42 @@ def main(): print("[+] Build directory removed") return - is_debug = args.debug or args.sanitize or args.coverage + # ENABLE_SANITIZERS only takes effect inside the CMake tests block, so without --test the + # flag is accepted and silently does nothing, leaving an unsanitized build in build/sanitize. + # Refuse here rather than after a Conan install, so no build time is spent on it. + if sum((args.sanitize, args.tsan, args.msan)) > 1: + parser.error("--sanitize, --tsan and --msan cannot be combined: each sanitizer needs a " + "build of its own") + if args.msan and args.coverage: + parser.error("--msan and --coverage cannot be combined") + if args.tsan and not args.test: + parser.error( + "--tsan requires --test: the sanitizer flags are only applied to a build with " + "tests, so this would produce an unsanitized build in build/tsan" + ) + if args.sanitize and not args.test: + parser.error( + "--sanitize requires --test: the sanitizer flags are only applied to a build with " + "tests, so this would produce an unsanitized build in build/sanitize" + ) + + is_debug = args.debug or args.sanitize or args.tsan or args.msan or args.coverage build_type = "Debug" if is_debug else "Release" if args.sanitize: build_name = "sanitize" + elif args.tsan: + build_name = "tsan" + elif args.msan: + build_name = "msan" elif args.coverage: build_name = "coverage" elif is_debug: build_name = "debug" else: build_name = "release" + if args.compiler == "clang" and not args.msan: + build_name += "-clang" if args.cppstd != "20": build_name += f"-cxx{args.cppstd}" build_dir = f"build/{build_name}" @@ -225,17 +384,32 @@ def main(): f"-DBUILD_TESTING={'ON' if args.test else 'OFF'}", f"-DBUILD_EXAMPLES={'ON' if args.examples else 'OFF'}", f"-DBUILD_DOCS={'ON' if args.docs else 'OFF'}", + f"-DMCP_CPP_SDK_BUILD_CONFORMANCE={'ON' if args.conformance else 'OFF'}", f"-DMCP_CPP_SDK_BUILD_SHARED={'ON' if args.linkage in ('both', 'shared') else 'OFF'}", f"-DMCP_CPP_SDK_BUILD_STATIC={'ON' if args.linkage in ('both', 'static') else 'OFF'}", f"-DMCP_CPP_SDK_DEFAULT_LINKAGE={'static' if args.linkage == 'static' else 'shared'}", ] + clang = None + if args.compiler == "clang" or args.msan: + clang = find_clang() + if clang is None: + parser.error(f"this build needs clang and clang++ (or clang-{CLANG_VERSION} and " + f"clang++-{CLANG_VERSION}) on PATH") + conan_args = [] + if args.msan: + # The Conan toolchain file names the compiler and carries the flags. + conan_args = msan_conan_args(build_msan_libcxx(MSAN_LIBCXX_DIR, clang, args.jobs), clang) + elif clang: + extra_cmake += [f"-DCMAKE_C_COMPILER={clang[0]}", f"-DCMAKE_CXX_COMPILER={clang[1]}"] if args.sanitize: extra_cmake.append("-DENABLE_SANITIZERS=ON") + if args.tsan: + extra_cmake.append("-DENABLE_TSAN=ON") if args.coverage: extra_cmake.append("-DENABLE_COVERAGE=ON") generators_dir = conan_install( - build_dir, args.jobs, args.cppstd, build_type + build_dir, args.jobs, args.cppstd, build_type, conan_args ) cmake_configure( build_dir, @@ -243,18 +417,29 @@ def main(): generators_dir / "conan_toolchain.cmake", *extra_cmake, ) - cmake_build(build_dir, args.jobs) + aslr_prefix = no_aslr() if args.tsan or args.msan else [] + cmake_build(build_dir, args.jobs, aslr_prefix) write_user_presets(generators_dir) if args.test: test_jobs = args.jobs extra_env = ( - {"ASAN_OPTIONS": "detect_leaks=0", + {"ASAN_OPTIONS": "detect_leaks=1:detect_stack_use_after_return=1:strict_string_checks=1" + ":detect_invalid_pointer_pairs=2", "UBSAN_OPTIONS": "print_stacktrace=1:halt_on_error=1"} if args.sanitize else None ) - run("ctest", "--test-dir", build_dir, f"-j{test_jobs}", "--output-on-failure", - extra_env=extra_env) + if args.tsan: + suppressions = Path("test/tsan.supp").resolve() + extra_env = {"TSAN_OPTIONS": f"halt_on_error=1:second_deadlock_stack=1:suppressions={suppressions}"} + ctest_args = [] + if args.msan: + extra_env = {"MSAN_OPTIONS": "halt_on_error=1"} + # This test compiles a consumer with the compiler's own standard library, which + # cannot link against an SDK built on the instrumented libc++. + ctest_args = ["-E", "^packaging-pkgconfig-static-consumer$"] + run(*aslr_prefix, "ctest", "--test-dir", build_dir, f"-j{test_jobs}", "--output-on-failure", + *ctest_args, extra_env=extra_env) if args.coverage: run("gcovr", "-r", ".", "--html", "--html-details", diff --git a/scripts/check_alpha_runner_guard.py b/scripts/check_alpha_runner_guard.py new file mode 100644 index 0000000..caf9b66 --- /dev/null +++ b/scripts/check_alpha_runner_guard.py @@ -0,0 +1,115 @@ +#!/usr/bin/env python3 +"""Guard checks that keep the alpha conformance runner out of CI gating. + +Keeps the alpha conformance runner side-channel ( +``@modelcontextprotocol/conformance@0.2.0-alpha.11``, spec ``2026-07-28``) +permanently non-default and non-CI-gating: + + (i) ``conformance/run.sh`` must stay pinned to the locked ``0.1.16`` + runner used for tier evidence. + (ii) no file under ``.github/workflows/`` or ``scripts/`` (this guard + excluded) may reference the alpha results directory + (``conformance-results-alpha``) as a tier-evidence path. + (iii) no file under ``.github/workflows/`` or ``scripts/`` (this guard + excluded) may reference the alpha runner entrypoint + (``run_alpha.sh``) as something it executes and gates on. + +This intentionally does not require ``conformance/run_alpha.sh`` to exist: +it only asserts that CI-facing workflows and phase-exit scripts do not +*consume* the alpha runner as a pass/fail condition. Pattern-based, not +presence-based. +""" + +from __future__ import annotations + +import re +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parent.parent +RUN_SH = REPO_ROOT / "conformance" / "run.sh" +PINNED_VERSION = "0.1.16" + +SCAN_DIRS = ( + REPO_ROOT / ".github" / "workflows", + REPO_ROOT / "scripts", +) + +# This guard must itself name the forbidden patterns, so it excludes its own +# file from the scan rather than tripping over its own docstring/regex. +SELF_PATH = Path(__file__).resolve() + +FORBIDDEN_PATTERN = re.compile(r"conformance-results-alpha|run_alpha\.sh") + + +class GuardError(RuntimeError): + """Raised when an alpha-runner non-gating invariant is violated.""" + + +def check_pin() -> None: + if not RUN_SH.is_file(): + raise GuardError(f"{RUN_SH} does not exist; cannot verify the {PINNED_VERSION} pin") + text = RUN_SH.read_text(encoding="utf-8") + if PINNED_VERSION not in text: + raise GuardError( + f"{RUN_SH.relative_to(REPO_ROOT)} no longer contains the pinned runner version " + f"'{PINNED_VERSION}' (run.sh stays pinned while the alpha " + "harness is a separate, non-default, non-CI-gating side channel)" + ) + + +def iter_scanned_files() -> list[Path]: + files: list[Path] = [] + for scan_dir in SCAN_DIRS: + if not scan_dir.is_dir(): + continue + for path in sorted(scan_dir.rglob("*")): + if not path.is_file(): + continue + if path.resolve() == SELF_PATH: + continue + files.append(path) + return files + + +def check_no_gating_reference() -> None: + violations: list[str] = [] + for path in iter_scanned_files(): + try: + text = path.read_text(encoding="utf-8") + except (UnicodeDecodeError, OSError): + continue + for lineno, line in enumerate(text.splitlines(), start=1): + if FORBIDDEN_PATTERN.search(line): + violations.append(f"{path.relative_to(REPO_ROOT)}:{lineno}: {line.strip()!r}") + + if violations: + joined = "\n ".join(violations) + raise GuardError( + "the alpha runner/results directory is referenced by a CI workflow or " + "scripts/ guard, which would let it gate CI:\n " + f"{joined}\n" + "The alpha runner must remain non-default and non-CI-gating. " + "Remove the reference from any file " + "under .github/workflows/ or scripts/; the alpha harness itself " + "(conformance/run_alpha.sh) legitimately owns these strings and is not scanned." + ) + + +def main() -> int: + check_pin() + check_no_gating_reference() + print( + f"alpha-runner guard passed: {RUN_SH.relative_to(REPO_ROOT)} pinned to " + f"{PINNED_VERSION}; no CI workflow or scripts/ guard references the alpha " + "runner/results as a gating condition" + ) + return 0 + + +if __name__ == "__main__": + try: + raise SystemExit(main()) + except GuardError as error: + print(f"alpha-runner guard failed: {error}", file=sys.stderr) + raise SystemExit(1) diff --git a/scripts/check_json_matrix.py b/scripts/check_json_matrix.py new file mode 100644 index 0000000..acc1b9b --- /dev/null +++ b/scripts/check_json_matrix.py @@ -0,0 +1,199 @@ +#!/usr/bin/env python3 +"""Fail the build when the peer-input matrix has fallen behind the protocol. + +The matrix's value is not its size. A hand-written suite can be just as large +and still miss a type nobody thought about -- which is what happened three +times to the same defect class. What makes a generated matrix different is +that its type list is *checked* against the protocol, so a new type cannot be +added without either entering the matrix or being excluded on the record. + +This script enforces exactly that: + + { types with a from_json under include/mcp/protocol/ } + == { types in the matrix } u { types excluded, with a stated reason } + +It also verifies that the generated file is current. The checked-in matrix +and manifest must equal what ``gen_json_matrix.render()`` produces from the +current sources, compared with all whitespace removed so clang-format's line +wrapping does not count. On a mismatch the rows on both sides are parsed to +say which way each one moved: + + stale a decoder fix landed, validation got stricter, rows or types + came or went, or the harness text was edited -- regenerate + (exit 1) + REGRESSION a member now throws on an explicit null while it may be + absent, or has become required, where the committed matrix + says it tolerated the input -- fix the decoder (exit 2) + +Regenerating over a regression would record it as expected behaviour, so +``gen_json_matrix.py`` refuses to unless the row is named with +``--accept-regression``. +""" + +from __future__ import annotations + +import argparse +import json +import os +import sys + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + +import gen_json_matrix as gen # noqa: E402 + + +MATRIX = "test/core/json_peer_input_matrix_test.cpp" +REGENERATE = " python3 scripts/gen_json_matrix.py" + + +def check( + committed_text: str, + committed_manifest: dict, + rendered_text: str, + rendered_manifest: dict, +) -> tuple[int, list[str]]: + """Compares the committed matrix with a fresh render; returns (exit code, lines).""" + if ( + gen.strip_ws(committed_text) == gen.strip_ws(rendered_text) + and committed_manifest == rendered_manifest + ): + m = committed_manifest + return 0, [ + f"peer-input matrix: {len(m['types'])} types covered, {len(m['excluded'])} " + f"excluded on the record, {m['case_count']} cases " + f"({m['field_count']} fields x {m['mode_count']} modes); " + "matches scripts/gen_json_matrix.py" + ] + + committed = gen.parse_matrix(committed_text) + if not committed["tests"]: + return 1, [ + f"{MATRIX} has no TEST(JsonPeerInputMatrix, ...) to compare. Restore it " + "from git before regenerating, so the new matrix is compared with the old:", + f" git checkout -- {MATRIX}", + ] + regressions, stale = gen.classify(committed, gen.parse_matrix(rendered_text)) + # Text that differs with no difference the parse can name is harness text. + if not (regressions or stale) and gen.strip_ws(committed_text) != gen.strip_ws(rendered_text): + stale.append(gen.HARNESS_DIFFERS) + for key in sorted(set(committed_manifest) | set(rendered_manifest)): + was, now = committed_manifest.get(key), rendered_manifest.get(key) + if was == now: + continue + if key not in rendered_manifest: + stale.append(f"manifest key '{key}': only in the committed file") + elif key not in committed_manifest: + stale.append(f"manifest key '{key}': only in the generated file") + elif isinstance(was, (list, dict)) or isinstance(now, (list, dict)): + stale.append(f"manifest key '{key}': committed and generated differ") + else: + stale.append(f"manifest key '{key}': committed {was}, generated {now}") + + lines: list[str] = [] + if regressions: + lines += gen.regression_block(regressions) + if stale: + lines.append( + f"peer-input matrix is stale: {MATRIX} does not match what " + "scripts/gen_json_matrix.py generates from the current sources." + ) + lines += [f" {line}" for line in stale] + if not regressions and any(line.endswith("(now tolerated)") for line in stale): + lines += [ + "A row moving toward tolerance means a decoder fix landed. " + "Regenerate and commit the result with the fix:", + REGENERATE, + ] + elif not regressions: + lines += [ + "The generated matrix changed (template, rows added or removed, or stricter " + "validation). Review the diff, then regenerate:", + REGENERATE, + ] + return (2 if regressions else 1), lines + + +def main(argv: list[str]) -> int: + ap = argparse.ArgumentParser(description=__doc__) + repo = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + ap.add_argument("--repo", default=repo) + ap.add_argument( + "--manifest", + default=os.path.join(repo, "test", "core", "json_matrix_manifest.json"), + ) + ap.add_argument( + "--matrix", + default=os.path.join(repo, "test", "core", "json_peer_input_matrix_test.cpp"), + ) + args = ap.parse_args(argv) + + if not os.path.exists(args.manifest): + sys.stderr.write( + f"{args.manifest} is missing; run: python3 scripts/gen_json_matrix.py\n" + ) + return 1 + + with open(args.manifest, encoding="utf-8") as fh: + manifest = json.load(fh) + _structs, _bases, _enums, _defaults, protocol_types = gen.load_protocol(args.repo) + + covered = set(manifest["types"]) + excluded = set(manifest["excluded"]) + accounted = covered | excluded + + missing = sorted(protocol_types - accounted) + stale = sorted(accounted - protocol_types) + + failures = 0 + if missing: + failures += 1 + sys.stderr.write( + "These protocol types have a from_json but are not in the peer-input\n" + "matrix and are not excluded:\n" + ) + for name in missing: + sys.stderr.write(f" {name}\n") + sys.stderr.write( + "\nRegenerate the matrix (python3 scripts/gen_json_matrix.py). If a type\n" + "genuinely cannot be exercised, add it to UNSYNTHESISABLE in\n" + "scripts/gen_json_matrix.py with the reason -- an exclusion on the\n" + "record is fine, an omission nobody noticed is what this check exists\n" + "to prevent.\n" + ) + if stale: + failures += 1 + sys.stderr.write( + "These types are in the matrix manifest but no longer have a from_json\n" + "under include/mcp/protocol/:\n" + ) + for name in stale: + sys.stderr.write(f" {name}\n") + + expected_cases = manifest["field_count"] * manifest["mode_count"] + if manifest["case_count"] != expected_cases: + failures += 1 + sys.stderr.write( + f"case count {manifest['case_count']} is not fields x modes " + f"({manifest['field_count']} x {manifest['mode_count']} = {expected_cases})\n" + ) + + if failures: + return 1 + + if not os.path.exists(args.matrix): + sys.stderr.write( + f"{args.matrix} is missing. Restore it from git before regenerating, so " + f"the new matrix is compared with the old:\n git checkout -- {MATRIX}\n" + ) + return 1 + with open(args.matrix, encoding="utf-8") as fh: + committed_text = fh.read() + rendered_text, rendered_manifest = gen.render(args.repo) + code, lines = check(committed_text, manifest, rendered_text, rendered_manifest) + for line in lines: + sys.stderr.write(line + "\n") + return code + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) diff --git a/scripts/check_readme_snippets.py b/scripts/check_readme_snippets.py new file mode 100644 index 0000000..fd745ab --- /dev/null +++ b/scripts/check_readme_snippets.py @@ -0,0 +1,187 @@ +#!/usr/bin/env python3 +""" +Compile the complete C++ snippets embedded in README.md. + +The README is the first thing a new user copies, so it is compiled the same way +a user would compile it: each fenced ``cpp`` block that contains ``int main(`` +is written out as a standalone translation unit and fed to the compiler that +built the SDK, with the SDK's own include paths and language standard. + +Fenced blocks without ``int main(`` are fragments that only make sense inside a +larger program, so they are reported as skipped rather than compiled. + +Compiler flags are taken from ``compile_commands.json`` in the build directory, +so this script needs a configured build tree. With no --build-dir it looks for +one under build/, preferring build/release, which is what scripts/build.py +produces. +""" +import argparse +import json +import re +import shlex +import subprocess +import sys +import tempfile +from pathlib import Path + +FENCE_RE = re.compile(r"^```cpp\s*$(.*?)^```\s*$", re.MULTILINE | re.DOTALL) + +# Flags that carry the information a snippet needs: where the headers are, which +# language standard applies, which macros the SDK's own headers expect, and which +# toolchain root the compiler was aimed at. The last group is what makes this work +# on macOS: CMake drives an Xcode clang with ``-isysroot ``, and without it +# that compiler cannot find its own ````, so every snippet fails for a +# reason that has nothing to do with the snippet. +PASSTHROUGH_PREFIXES = ("-I", "-D", "-std=", "-m", "-stdlib=", "--sysroot=", + "--target=") +PASSTHROUGH_PAIRS = ("-isystem", "-include", "-imacros", "-isysroot", "--sysroot", + "-arch", "-target") + + +def extract_snippets(readme_path): + text = readme_path.read_text(encoding="utf-8") + snippets = [] + for match in FENCE_RE.finditer(text): + body = match.group(1) + line_no = text.count("\n", 0, match.start()) + 1 + snippets.append((line_no, body)) + return snippets + + +def has_example_entry(database): + try: + entries = json.loads(database.read_text(encoding="utf-8")) + except (OSError, ValueError): + return False + return any("/examples/" in item["file"].replace("\\", "/") for item in entries) + + +def find_build_dir(explicit): + if explicit: + return Path(explicit) + preferred = Path("build/release") + candidates = [preferred / "compile_commands.json"] + candidates += sorted(Path("build").glob("*/compile_commands.json")) + for database in candidates: + if database.is_file() and has_example_entry(database): + return database.parent + return preferred + + +def compile_flags(build_dir): + database = Path(build_dir) / "compile_commands.json" + if not database.is_file(): + raise FileNotFoundError( + f"{database} not found. Configure a build first, " + f"for example: python scripts/build.py --examples" + ) + + entries = json.loads(database.read_text(encoding="utf-8")) + # Only an example's compile line will do. A snippet is consumer code, and the + # library's own translation units are compiled with export-side macros that a + # consumer must never see. + entry = next( + (item for item in entries if "/examples/" in item["file"].replace("\\", "/")), + None, + ) + if entry is None: + raise RuntimeError( + f"{database} holds no compile command for an example, so no consumer-side " + f"flags are available. Build with examples first, " + f"for example: python scripts/build.py --examples" + ) + + argv = entry.get("arguments") or shlex.split(entry["command"]) + compiler = argv[0] + + flags = [] + index = 1 + while index < len(argv): + arg = argv[index] + if arg in PASSTHROUGH_PAIRS and index + 1 < len(argv): + flags.extend([arg, argv[index + 1]]) + index += 2 + continue + if arg.startswith(PASSTHROUGH_PREFIXES): + flags.append(arg) + index += 1 + + return compiler, flags + + +def main(): + parser = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--build-dir", default=None, + help="Configured build directory holding compile_commands.json " + "(default: build/release, or the first build/* that has one)") + parser.add_argument("--readme", default="README.md", help="Markdown file to check") + args = parser.parse_args() + + readme_path = Path(args.readme) + snippets = extract_snippets(readme_path) + if not snippets: + print(f"No cpp snippets found in {readme_path}") + return 1 + + build_dir = find_build_dir(args.build_dir) + try: + compiler, flags = compile_flags(build_dir) + except (FileNotFoundError, RuntimeError) as error: + print(f"error: {error}") + return 1 + print(f"Build directory: {build_dir}") + print(f"Compiler: {compiler}") + print(f"Flags: {' '.join(flags)}\n") + + compiled = 0 + skipped = 0 + failed = 0 + + with tempfile.TemporaryDirectory() as tmp_dir: + for line_no, body in snippets: + label = f"{readme_path}:{line_no}" + if "int main(" not in body: + print(f"{label}: SKIP (fragment, no main)") + skipped += 1 + continue + + source = Path(tmp_dir) / f"readme_{line_no}.cpp" + source.write_text(body, encoding="utf-8") + + print(f"{label}: compiling...", end=" ", flush=True) + result = subprocess.run( + [compiler, "-fsyntax-only", *flags, str(source)], + capture_output=True, + text=True, + ) + if result.returncode == 0: + print("PASS") + compiled += 1 + else: + print("FAIL") + print("--- snippet ---") + print(body.rstrip()) + if result.stdout: + print("--- stdout ---") + print(result.stdout.rstrip()) + if result.stderr: + print("--- stderr ---") + print(result.stderr.rstrip()) + failed += 1 + + print(f"\n{'=' * 50}") + print(f"Results: {compiled} compiled, {skipped} skipped, {failed} failed") + + if failed: + print("\nSome README snippets do not compile!") + return 1 + if compiled == 0: + print("\nNo complete README snippet was compiled; the check proved nothing.") + return 1 + print("\nAll README snippets compiled.") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/check_release_workflow.py b/scripts/check_release_workflow.py index 922e0e7..bdc5adb 100644 --- a/scripts/check_release_workflow.py +++ b/scripts/check_release_workflow.py @@ -68,6 +68,7 @@ "secrets.RELEASE_GPG_PRIVATE_KEY_B64", "secrets.RELEASE_GPG_PASSPHRASE", "secrets.AUR_SSH_PRIVATE_KEY_B64", + "secrets.AUR_SSH_KEY_PASSPHRASE", "secrets.HOMEBREW_APP_PRIVATE_KEY", "secrets.CHOCOLATEY_APP_PRIVATE_KEY", "secrets.CONAN_BROKER_APP_PRIVATE_KEY", @@ -287,6 +288,19 @@ def check_publication_invariants(blocks: dict[str, list[str]], text: str) -> Non block = "\n".join(blocks[job]) if "github-anchor" not in block or "Verify fixed anchor handoff" not in block: raise PolicyError(f"publisher {job!r} is not bound to the verified GitHub anchor") + aur = "\n".join(blocks["aur"]) + for fragment in ( + "secrets.AUR_SSH_KEY_PASSPHRASE", + "SSH_ASKPASS_REQUIRE=force", + "ssh-add", + "BatchMode=yes", + "StrictHostKeyChecking=yes", + "UserKnownHostsFile=${HOME}/.ssh/known_hosts", + ): + if fragment not in aur: + raise PolicyError(f"AUR publisher is missing passphrase-protected SSH handling {fragment!r}") + if "ssh-keyscan" in aur or "StrictHostKeyChecking=no" in aur: + raise PolicyError("AUR publisher weakens SSH host-key verification") cloudsmith = "\n".join(blocks["publish-apt"] + blocks["publish-rpm"]) if "cloudsmith-cli-action@159f1619275d5d3147f059c3cc110938ec221d16" not in cloudsmith: raise PolicyError("Cloudsmith action pin is missing") diff --git a/scripts/check_stdio_streams.py b/scripts/check_stdio_streams.py new file mode 100644 index 0000000..feeb6d2 --- /dev/null +++ b/scripts/check_stdio_streams.py @@ -0,0 +1,155 @@ +#!/usr/bin/env python3 +""" +Assert that the stdio examples keep stdout clean for the protocol. + +On the stdio transport stdout *is* the JSON-RPC channel, so a human-readable +line printed there corrupts the message stream for the peer. Running an example +and checking its exit code does not notice this: the process still exits 0 while +emitting garbage the peer cannot parse. + +This check drives each stdio example through a short scripted exchange and +requires every non-empty line it writes to stdout to parse as a JSON-RPC 2.0 +message. Diagnostics are expected on stderr and are not inspected. +""" +import json +import subprocess +import sys +import time +from pathlib import Path + +CLIENT_INITIALIZE_RESPONSE = json.dumps({ + "jsonrpc": "2.0", + "id": "1", + "result": { + "protocolVersion": "2025-11-25", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "stream-check-server", "version": "1.0.0"}, + }, +}) + +SERVER_REQUESTS = [ + json.dumps({ + "jsonrpc": "2.0", + "id": 1, + "method": "initialize", + "params": { + "protocolVersion": "2025-06-18", + "clientInfo": {"name": "stream-check", "version": "1.0.0"}, + "capabilities": {}, + }, + }), + json.dumps({"jsonrpc": "2.0", "method": "notifications/initialized"}), + json.dumps({"jsonrpc": "2.0", "id": 2, "method": "tools/list"}), +] + +# Each case is (binary name, lines fed to stdin). A binary that is not present +# in the build directory is skipped, so this runs against partial builds. +CASES = [ + ("example-client-stdio", [CLIENT_INITIALIZE_RESPONSE]), + ("example-server-stdio", SERVER_REQUESTS), +] + + +def executable(build_dir, name): + for candidate in (build_dir / name, build_dir / f"{name}.exe"): + if candidate.is_file(): + return candidate + return None + + +def run_exchange(path, stdin_lines, settle_seconds=3.0): + """Feed the example its scripted input, then hold stdin open. + + Closing stdin straight away makes the peer see EOF and exit before it has + printed anything, which would let a corrupting example pass. Keeping the + pipe open lets the exchange run far enough for diagnostics to appear. + """ + process = subprocess.Popen( + [str(path)], + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + ) + try: + for line in stdin_lines: + process.stdin.write(line + "\n") + process.stdin.flush() + time.sleep(settle_seconds) + except (BrokenPipeError, OSError): + pass + + try: + process.stdin.close() + except (BrokenPipeError, OSError): + pass + # communicate() would try to flush the pipe we just closed. + process.stdin = None + + try: + stdout, _ = process.communicate(timeout=15) + except subprocess.TimeoutExpired: + process.kill() + stdout, _ = process.communicate() + return stdout + + +def check_stream(name, path, stdin_lines): + stdout = run_exchange(path, stdin_lines) + + offenders = [] + messages = 0 + for number, line in enumerate(stdout.splitlines(), start=1): + if not line.strip(): + continue + try: + message = json.loads(line) + except ValueError: + offenders.append((number, line)) + continue + if not isinstance(message, dict) or message.get("jsonrpc") != "2.0": + offenders.append((number, line)) + continue + messages += 1 + + if offenders: + print(f"{name}: FAIL") + print(" stdout is the JSON-RPC channel, but these lines are not JSON-RPC:") + for number, line in offenders[:10]: + print(f" line {number}: {line[:160]}") + print(" Send diagnostics to stderr instead.") + return False + + if messages == 0: + print(f"{name}: FAIL (no JSON-RPC written to stdout; the check proved nothing)") + return False + + print(f"{name}: PASS ({messages} JSON-RPC messages, no stray output)") + return True + + +def main(): + build_dir = Path(sys.argv[1]) if len(sys.argv) > 1 else Path("build/release") + + checked = 0 + failed = 0 + for name, stdin_lines in CASES: + path = executable(build_dir, name) + if path is None: + print(f"{name}: SKIP (not built in {build_dir})") + continue + checked += 1 + if not check_stream(name, path, stdin_lines): + failed += 1 + + print(f"\n{'=' * 50}") + if checked == 0: + print(f"No stdio examples found in {build_dir}; build them with " + f"python scripts/build.py --examples") + return 1 + print(f"Results: {checked - failed} clean, {failed} corrupting stdout") + return 1 if failed else 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/gen_json_matrix.py b/scripts/gen_json_matrix.py new file mode 100644 index 0000000..8d84d51 --- /dev/null +++ b/scripts/gen_json_matrix.py @@ -0,0 +1,967 @@ +#!/usr/bin/env python3 +"""Generate the peer-input test matrix from the census. + +For every type a peer can steer a document into, for every field of that type, +the matrix exercises four modes: + + absent the key is not in the document + null the key is present and explicitly null + wrong_type the key holds a value of the wrong JSON type + oversized the key holds a very large value of the right JSON type + +Each case asserts what ``scripts/json_census.py`` predicts for that field and +mode. The suite is therefore a differential oracle rather than a restatement: +where the code and the census disagree, one of them is wrong and the build says +so. A hand-written suite of the same size asserts only what its author already +believed, which is how the same defect class survived three reviews. + +The generated file also pins the matrix's type list. ``check_json_matrix.py`` +compares that list against the types with a ``from_json`` under +``include/mcp/protocol/`` and fails the build when a new protocol type has not +been added, so the matrix cannot silently fall behind the protocol. + +Regenerating cannot launder a regression either. When a row would move from +tolerating an explicit null or an absent key to throwing on it, nothing is +written unless that row is named with ``--accept-regression Type.key`` (or a +bare ``Type``), and the reason belongs in the commit message. Rows are judged +against both the matrix in the working tree and the one committed at git HEAD, +so editing or deleting the working copy first does not hide a regression. +Only ``--no-baseline``, for a tree where neither exists, writes unchecked. +""" + +from __future__ import annotations + +import argparse +import json +import os +import re +import shutil +import subprocess +import sys + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + +import json_census as census # noqa: E402 + +PROTOCOL_DIR = os.path.join("include", "mcp", "protocol") +MODES = ("absent", "null", "wrong_type", "oversized") + +# Enum serialisers give the census a valid string for an enum-typed field. +ENUM_MACRO = re.compile( + r"NLOHMANN_JSON_SERIALIZE_ENUM\s*\(\s*([\w:]+)\s*,\s*\{(.*?)\}\s*\)\s*$", + re.DOTALL | re.MULTILINE, +) + +SCALARS = { + "bool": True, + "int": 1, + "int64_t": 1, + "std::int64_t": 1, + "uint64_t": 1, + "size_t": 1, + "double": 1.0, + "float": 1.0, + "std::string": "x", + "std::string_view": "x", +} + +# Types with a hand-written from_json that dispatches on the document's shape +# rather than reading a fixed field set. A synthesised baseline cannot pick +# the right arm, so they are excluded by name, with the reason recorded here +# and reprinted in the generated file. check_json_matrix.py counts them as +# covered so the build stays honest about what the matrix does not reach. +UNSYNTHESISABLE = { + "JSONRPCResponse": "std::variant; from_json dispatches on which key is present", + "JSONRPCMessage": "std::variant; from_json dispatches on which key is present", + "ContentBlock": "std::variant; from_json dispatches on the 'type' discriminator", + "ResourceContents": "std::variant; from_json dispatches on text vs blob", + "CompleteReference": "std::variant; from_json dispatches on the 'type' discriminator", + "ElicitRequestParams": "std::variant; from_json dispatches on mode", + "SamplingMessageContent": "std::variant; from_json accepts a block or a list of blocks", + "PrimitiveSchemaDefinition": "base of EnumSchema; exercised through its derived type", + "RequestId": "std::variant of string and integer; it has no fields to vary", +} + + +def enum_values(text: str) -> dict[str, str]: + """Enum type -> a valid wire string for it.""" + out: dict[str, str] = {} + for m in ENUM_MACRO.finditer(text): + name = m.group(1).split("::")[-1] + first = re.search(r'"([^"]+)"', m.group(2)) + if first: + out[name] = first.group(1) + return out + + +MEMBER_WITH_DEFAULT = re.compile( + r"^\s*(?:const\s+)?[\w:]+(?:\s*<[^;]*>)?\s+(\w+)\s*=\s*([^;]+);\s*$", + re.MULTILINE, +) + + +def member_defaults(text: str, root) -> dict[tuple[str, str], object]: + """(type, member) -> the literal the member is declared with. + + A member declared `std::string type = "audio";` is a discriminator, and + `std::string jsonrpc = "2.0";` is checked by a validator. Synthesising + "x" for either produces a baseline the serialiser rejects outright, so + every mode of every field of that type fails for a reason that has nothing + to do with the mode. The declared default is the value the type expects. + """ + out: dict[tuple[str, str], object] = {} + + def walk(blk) -> None: + m = census.STRUCT_DECL.search(blk.header + "{") + if m: + flat = census.strip_nested(text[blk.start : blk.end]) + for member, literal in MEMBER_WITH_DEFAULT.findall(flat): + lit = literal.strip() + if lit.startswith('"') and lit.endswith('"'): + out[(m.group(1), member)] = lit[1:-1] + elif re.fullmatch(r"-?\d+", lit): + out[(m.group(1), member)] = int(lit) + elif lit in ("true", "false"): + out[(m.group(1), member)] = lit == "true" + for child in blk.children: + walk(child) + + walk(root) + return out + + +def load_protocol(base: str): + """Struct members, enum values, from_json targets, and declared defaults.""" + structs: dict[str, list[tuple[str, str]]] = {} + bases: dict[str, list[str]] = {} + enums: dict[str, str] = {} + defaults: dict[tuple[str, str], object] = {} + from_json_types: set[str] = set() + top = os.path.join(base, PROTOCOL_DIR) + for name in sorted(os.listdir(top)): + if not name.endswith(".hpp"): + continue + with open(os.path.join(top, name), encoding="utf-8") as fh: + raw = fh.read() + text = census.strip_comments(raw) + root = census.parse_blocks(text) + for k, v in census.parse_structs(text, root).items(): + structs.setdefault(k, v) + for k, v in census.parse_bases(text, root).items(): + bases.setdefault(k, v) + enums.update(enum_values(text)) + defaults.update(member_defaults(text, root)) + # Keep the qualification: the census keys a nested capability by + # `ClientCapabilities::RootsCapability`, and shortening it here left + # twenty types looking unsynthesisable when they were simply not being + # looked up under the name they were filed under. + for m in census.FROM_JSON_SIG.finditer(text): + from_json_types.add(m.group(2)) + for m in census.MACRO_DEFINE.finditer(text): + from_json_types.add(m.group(3)) + return structs, bases, enums, defaults, from_json_types + + +def unwrap(ctype: str) -> tuple[str, str]: + """Strip one layer of optional/vector/map. Returns (wrapper, inner).""" + c = ctype.strip().removeprefix("const ").strip() + for wrapper in ("std::optional", "std::vector", "std::map", "std::unordered_map"): + m = re.match(rf"{re.escape(wrapper)}\s*<(.+)>$", c) + if m: + return wrapper, m.group(1).strip() + return "", c + + +# Set once per run by render(); sample() needs the census view to build a +# nested type's baseline and is called from helpers that do not carry it. +_GOVERNED: dict = {} +_DEFAULTS: dict = {} +_BASES: dict = {} + + +def inherited_fields(type_name: str, governed, depth: int = 0): + """This type's fields plus every field its bases contribute.""" + seen = [(k, r) for (t, k), r in governed.items() if t == type_name] + if depth < 4: + for parent in _BASES.get(type_name.split("::")[-1], []): + have = {k for k, _ in seen} + seen += [ + (k, r) for k, r in inherited_fields(parent, governed, depth + 1) if k not in have + ] + return seen + + +def sample(ctype: str, structs, enums, depth: int = 0): + """A JSON value that deserialises into ``ctype``, or None if unreachable.""" + return sample_for(ctype, structs, enums, _GOVERNED, depth) + + +def field_type(type_name: str, row, structs) -> str: + """The declared C++ type behind a wire key.""" + if row.member: + for ctype, mname in structs.get(type_name.split("::")[-1], []): + if mname == row.member: + return ctype + return ",".join(row.target_types) + + +def baseline(type_name: str, structs, enums, governed, depth: int = 0): + """A minimal document that deserialises into ``type_name``. + + Keys come from the census, not from member names. They differ: `Icon` + reads its `source` member from `"src"`. A baseline built from member + names omitted that key, and every case for the type then threw on a + missing `"src"` rather than on the mode it was meant to exercise -- 17 + cases that looked like findings and were nothing but a bad fixture. + """ + if depth > 4: + return None + keys = inherited_fields(type_name, governed) + if not keys: + return None + doc: dict = {} + for key, row in keys: + if row.absent != census.THROW: + continue + ctype = field_type(row.owner, row, structs) + value = _DEFAULTS.get((row.owner.split("::")[-1], row.member)) + if value is None: + value = sample_for(ctype, structs, enums, governed, depth + 1) + if value is None: + return None + doc[key] = value + return doc + + +def sample_for(ctype: str, structs, enums, governed, depth: int = 0): + """Like ``sample``, but recursing through census-derived baselines.""" + if depth > 4 or not ctype: + return None + wrapper, inner = unwrap(ctype) + if wrapper == "std::optional": + return sample_for(inner, structs, enums, governed, depth + 1) + if wrapper == "std::vector": + item = sample_for(inner, structs, enums, governed, depth + 1) + return [] if item is None else [item] + if wrapper in ("std::map", "std::unordered_map"): + return {} + bare = inner.split("::")[-1] + if inner in SCALARS: + return SCALARS[inner] + if bare in SCALARS: + return SCALARS[bare] + if "nlohmann::json" in inner or inner == "json": + return {} + if bare in enums: + return enums[bare] + if bare in ("RequestId", "ProgressToken"): + return 1 + return baseline(bare, structs, enums, governed, depth + 1) + + +# Mutations are applied in C++, not baked into the file. Spelling out a +# 64 KiB string and a 512-element array for each of 1268 cases produced a 17 MB +# source file -- a generated suite has to stay readable and reviewable, or it +# is just a binary blob that happens to compile. The generator emits the +# baseline and a per-field hint; the mutation helpers in the test do the rest. +WRONG_FROM_VALUE = 0 # derive the counter-example from the value's JSON type +WRONG_FORCE_OBJECT = 1 # for variant fields that accept both string and number + + +def wrong_hint(ctype: str) -> int: + """How the test should pick a value the field does not accept. + + A ``RequestId`` is a variant of string and integer, so "the other scalar + type" is still a value it accepts. Deriving the counter-example from the + sample would generate a case that passes for the wrong reason. + """ + _, inner = unwrap(ctype) + if inner.split("::")[-1] in ("RequestId", "ProgressToken"): + return WRONG_FORCE_OBJECT + return WRONG_FROM_VALUE + + +def field_rows(rows): + """(type, field) -> the merged verdict for that field. + + A field can be touched by several sites -- a validator that rejects a + shape, then the accessor that reads it. The field's verdict is the + strictest of them: if any site throws on a null, the peer sending a null + gets a throw. Taking the first row instead let a permissive accessor mask + a validator two lines above it. + """ + out: dict[tuple[str, str], census.Row] = {} + for r in rows: + if r.site_kind in ("delegate", "envelope_presence"): + continue + if not r.owner or r.owner.startswith("(") or r.field_name.startswith("<"): + continue + key = (r.owner, r.field_name) + prev = out.get(key) + if prev is None: + out[key] = r + continue + for axis in ("absent", "null", "wrong_type"): + if getattr(r, axis) == census.THROW: + setattr(prev, axis, census.THROW) + if prev.constrained == census.NA and r.constrained != census.NA: + prev.constrained = r.constrained + if not prev.member and r.member: + prev.member = r.member + if not prev.target_types and r.target_types: + prev.target_types = r.target_types + return out + + +def predict(row: census.Row, mode: str) -> bool: + """Does the census say this mode throws?""" + if mode == "absent": + return row.absent == census.THROW + if mode == "null": + return row.null == census.THROW + if mode == "wrong_type": + return row.wrong_type == census.THROW + # Oversized input is accepted -- no site in this codebase bounds a value's + # size -- unless the field has a value domain, in which case an arbitrary + # value of the right JSON type is outside it. That is a domain rejection, + # not a size limit, and the two must not be confused: reading it as a size + # limit would report a ceiling that does not exist. + return row.constrained != census.NA + + +def cpp_string(doc) -> str: + return json.dumps(doc, separators=(",", ":")) + + +HEADER = '''// Generated by scripts/gen_json_matrix.py from scripts/json_census.py. +// Do not edit by hand: regenerate with `python3 scripts/gen_json_matrix.py`. +// +// Every peer-deserialisable protocol type, every field, four modes: the key +// absent, the key present and explicitly null, the key holding the wrong JSON +// type, and the key holding an oversized value of the right type. +// +// Each case asserts what the census predicts, so a disagreement between the +// code and the census fails here rather than in a user\'s session. Serialising +// an absent optional as an explicit null is the default for Go\'s +// encoding/json without omitempty and for a naively dumped Python dataclass, +// which puts the null column within reach of a careless peer rather than only +// a hostile one. +// +// scripts/check_json_matrix.py runs as a build step. It fails when a protocol +// type has a from_json but is neither in this matrix nor excluded on the +// record, so the matrix cannot quietly fall behind the protocol, and when this +// file differs from what the generator produces from the current sources. +// +// WHAT THIS CORPUS IS, AND WHAT IT IS NOT +// +// It is a consistency oracle, not a correctness one. Every expectation is +// derived by static analysis of the guard construct at the decode site, and +// the test then runs the decoder. A case fails when runtime and static +// analysis disagree. It cannot fail because a field is wrong per the MCP +// spec. Nothing here is produced by executing the SDK, so the two sides are +// independent -- but agreement means consistency, not conformance. +// +// It follows that once a decode site is fixed and this file is regenerated, +// the matching cases pass BY CONSTRUCTION: the census re-reads the fixed +// source and predicts the new behaviour. Green here is not evidence that a +// field decodes correctly. That evidence lives in the hand-written, +// red-first tests in test/core/protocol_test.cpp and +// test/server/server_handlers_test.cpp. +// +// A TRUE null_throws ON AN OPTIONAL MEMBER RECORDS A DEFECT, NOT A SPEC +// +// The tuple (absent=false, null=true, wrong=true, oversized=false) describes +// a member that tolerates being absent but throws when a peer sends it as an +// explicit null. That is the defect class this project has repeatedly been +// caught by, and scripts/json_census_dispositions.json dispositions it "fix". +// {null_fragile} of the {field_total} field entries below still carry it; there were 87 before the +// explicit-null decoding fixes. Those rows pin behaviour as it is so the +// suite stays green. They do not endorse it. Fixing one makes the build check +// report this file stale until it is regenerated with the fix. The check reads +// the other direction as a REGRESSION: a row whose decoder now throws on an +// explicit null or on an absent key, where this file says it tolerates it. +// The generator refuses to write such a row unless it is named with +// --accept-regression. Never relax a fix to satisfy this file. +// +// This column is also blind to the other half of that class. A member that +// decodes an explicit null into an ENGAGED optional throws nothing, so it is +// recorded as null_throws = false whether or not it then re-encodes a member +// the peer never sent. _meta and annotations were corrupt in exactly that way +// and not one case below changed when they were fixed. Seeing that class +// needs a round-trip check, which this corpus does not have. + +#include + +#include + +#include + +#include + +namespace { + +using nlohmann::json; + +// How to pick a value a field does not accept. A variant of string and integer +// needs an object: "the other scalar type" is still something it accepts. +enum WrongHint { kWrongFromValue = 0, kWrongForceObject = 1 }; + +json wrong_typed(const json& value, int hint) { + if (hint == kWrongForceObject) { + return json{{"neither", "a string nor an integer"}}; + } + if (value.is_boolean()) { + return "not-a-bool"; + } + if (value.is_number()) { + return "not-a-number"; + } + if (value.is_string()) { + return 12345; + } + if (value.is_array()) { + return json{{"not", "an array"}}; + } + return json::array({"not an object"}); +} + +// A very large value that still has the type the field accepts. An oversized +// array stays an array *of its own element type*: filling it with strings +// would make it a wrong-type case wearing an oversized label, and the throw +// would then be read as a size limit that does not exist. +json oversized(const json& value) { + if (value.is_boolean()) { + return value; + } + if (value.is_number()) { + return json(1000000000000000000LL); + } + if (value.is_string()) { + return std::string(65536, \'A\'); + } + if (value.is_array()) { + if (value.empty()) { + return value; + } + json grown = json::array(); + for (int i = 0; i < 512; ++i) { + grown.push_back(value.front()); + } + return grown; + } + // An object grows by keys the serialiser ignores, so its shape is intact. + json grown = value.is_object() ? value : json::object(); + for (int i = 0; i < 256; ++i) { + grown["_pad" + std::to_string(i)] = std::string(256, \'A\'); + } + return grown; +} + +// Parses `doc` into T and reports whether that threw. +template +bool throws_on(const json& doc) { + try { + static_cast(doc.get()); + return false; + } catch (const std::exception&) { + return true; + } +} + +struct FieldExpect { + const char* key; + const char* present; // a value the field accepts, as JSON text + int wrong_hint; + bool absent_throws; + bool null_throws; + bool wrong_throws; + bool oversized_throws; +}; + +// Applies all four modes to one field and checks each against the census. +template +void check_field(const json& baseline, const FieldExpect& f) { + const json present = json::parse(f.present); + + json absent = baseline; + absent.erase(f.key); + { + SCOPED_TRACE("absent"); + EXPECT_EQ(throws_on(absent), f.absent_throws); + } + json with_null = baseline; + with_null[f.key] = nullptr; + { + SCOPED_TRACE("null"); + EXPECT_EQ(throws_on(with_null), f.null_throws); + } + json with_wrong = baseline; + with_wrong[f.key] = wrong_typed(present, f.wrong_hint); + { + SCOPED_TRACE("wrong_type"); + EXPECT_EQ(throws_on(with_wrong), f.wrong_throws); + } + json with_big = baseline; + with_big[f.key] = oversized(present); + { + SCOPED_TRACE("oversized"); + EXPECT_EQ(throws_on(with_big), f.oversized_throws); + } +} + +} // namespace +''' + + +def render(base: str) -> tuple[str, dict]: + """The matrix source and its manifest, unformatted and without writing.""" + rows = census.build_rows(base) + + structs, bases, enums, defaults, from_json_types = load_protocol(base) + global _DEFAULTS, _BASES + _DEFAULTS = defaults + _BASES = bases + governed = field_rows(rows) + global _GOVERNED + _GOVERNED = governed + + covered: list[tuple[str, dict, list]] = [] + skipped: dict[str, str] = dict(UNSYNTHESISABLE) + + for type_name in sorted(from_json_types): + if type_name in skipped: + continue + base_doc = baseline(type_name, structs, enums, governed) + if base_doc is None: + skipped[type_name] = "no baseline document could be synthesised" + continue + cases = [] + for key, row in sorted(inherited_fields(type_name, governed)): + ctype = field_type(row.owner, row, structs) + present = _DEFAULTS.get((row.owner.split("::")[-1], row.member)) + if present is None: + present = sample_for(ctype, structs, enums, governed) + if present is None: + continue + cases.append( + ( + key, + cpp_string(present), + wrong_hint(ctype), + predict(row, "absent"), + predict(row, "null"), + predict(row, "wrong_type"), + predict(row, "oversized"), + ) + ) + if not cases: + skipped[type_name] = "no census row governs any of its fields" + continue + covered.append((type_name, base_doc, cases)) + + null_fragile = sum(1 for _, _, c in covered for case in c if not case[3] and case[4]) + field_total = sum(len(c) for _, _, c in covered) + body = [ + HEADER.replace("{null_fragile}", str(null_fragile)).replace( + "{field_total}", str(field_total) + ) + ] + for type_name, base_doc, cases in covered: + test_name = type_name.replace("::", "_") + body.append(f"\nTEST(JsonPeerInputMatrix, {test_name}) {{") + body.append( + f' const json baseline = json::parse(R"json({cpp_string(base_doc)})json");' + ) + body.append(" static const FieldExpect fields[] = {") + for key, present, hint, a, n, w, o in cases: + body.append( + f' {{"{key}", R"json({present})json", {hint}, ' + f"{str(a).lower()}, {str(n).lower()}, {str(w).lower()}, {str(o).lower()}}}," + ) + body.append(" };") + body.append(" for (const auto& f : fields) {") + body.append(f' SCOPED_TRACE(std::string("{type_name}.") + f.key);') + body.append(f" check_field(baseline, f);") + body.append(" }") + body.append("}") + + total_cases = sum(len(c) for _, _, c in covered) * len(MODES) + body.append( + f""" +// The matrix's own arithmetic. `types x fields x modes` is what a generated +// suite is worth; a hand-written suite of the same size is only worth what its +// author thought to write down. +TEST(JsonPeerInputMatrix, CaseCountIsDerived) {{ + constexpr int types = {len(covered)}; + constexpr int fields = {total_cases // len(MODES)}; + constexpr int modes = {len(MODES)}; + constexpr int generated_cases = {total_cases}; + EXPECT_EQ(generated_cases, fields * modes); + EXPECT_GT(types, 0); +}} +""" + ) + + body.append("// Types in the matrix (checked against the protocol headers at build time):") + for type_name, _, _ in covered: + body.append(f"// {type_name}") + body.append("// Types deliberately not in the matrix, and why:") + for type_name, reason in sorted(skipped.items()): + body.append(f"// {type_name}: {reason}") + + manifest = { + "types": [t for t, _, _ in covered], + "excluded": skipped, + "case_count": total_cases, + "field_count": total_cases // len(MODES), + "mode_count": len(MODES), + "gtest_count": len(covered) + 1, + } + return "\n".join(body) + "\n", manifest + + +# The reader for the file render() writes. clang-format wraps some rows and +# baselines across lines, so every pattern allows whitespace anywhere, and all +# comparisons are on whitespace-stripped text. +_TEST_RE = re.compile(r"TEST\(\s*JsonPeerInputMatrix\s*,\s*(\w+)\s*\)\s*\{") +_BASELINE_RE = re.compile( + r'const\s+json\s+baseline\s*=\s*json::parse\(\s*R"json\((.*?)\)json"\s*\)\s*;', + re.DOTALL, +) +_ROW_RE = re.compile( + r'\{\s*"([^"]*)"\s*,\s*R"json\((.*?)\)json"\s*,\s*(\d+)\s*,' + r"\s*(true|false)\s*,\s*(true|false)\s*,\s*(true|false)\s*,\s*(true|false)\s*\}", + re.DOTALL, +) +_CONST_RE = re.compile(r"constexpr\s+int\s+(\w+)\s*=\s*(\d+)\s*;") +_TAIL_MARKER = "// Types in the matrix" +_HEADER_COUNT_RE = re.compile(r"\d+\s+of\s+the\s+\d+\s+field\s+entries") +COLUMNS = ("absent", "null", "wrong_type", "oversized") + +HARNESS_DIFFERS = ( + "harness text outside the case rows differs " + "(hand edit, or a template change in scripts/gen_json_matrix.py)" +) + + +def strip_ws(text: str) -> str: + return re.sub(r"\s+", "", text) + + +def parse_matrix(text: str) -> dict: + """Reads a matrix back into its tests, rows and harness, however wrapped.""" + cut = text.find(_TAIL_MARKER) + body, tail = (text, "") if cut < 0 else (text[:cut], text[cut:]) + heads = list(_TEST_RE.finditer(body)) + tests: dict[str, dict] = {} + for i, head in enumerate(heads): + end = heads[i + 1].start() if i + 1 < len(heads) else len(body) + block = body[head.start() : end] + base = _BASELINE_RE.search(block) + rows = {} + for r in _ROW_RE.finditer(block): + rows[r.group(1)] = { + "present": strip_ws(r.group(2)), + "hint": int(r.group(3)), + "absent": r.group(4) == "true", + "null": r.group(5) == "true", + "wrong_type": r.group(6) == "true", + "oversized": r.group(7) == "true", + } + # What is left once the generated values are taken out: the part of the + # block every TEST shares, so a hand edit to it is still seen. + skeleton = _ROW_RE.sub("", _BASELINE_RE.sub("", block)) + skeleton = _CONST_RE.sub(lambda c: c.group(1), skeleton) + tests[head.group(1)] = { + "baseline": strip_ws(base.group(1)) if base else None, + "rows": rows, + "consts": {c.group(1): int(c.group(2)) for c in _CONST_RE.finditer(block)}, + "skeleton": strip_ws(skeleton), + } + prefix = body[: heads[0].start()] if heads else body + # The header's null-fragile count follows from the rows, which are compared + # one by one; a hand edit to it still fails the whitespace-stripped match. + prefix = _HEADER_COUNT_RE.sub("of the field entries", prefix) + return {"prefix": strip_ws(prefix), "tests": tests, "tail": strip_ws(tail)} + + +def _null_fragile(row: dict) -> bool: + return not row["absent"] and row["null"] + + +def classify(committed: dict, generated: dict) -> tuple[list, list]: + """Sorts every difference by direction, committed -> generated. + + Returns (regressions, stale). A regression is a (test, key, line, kind) + for a member that now throws on input the committed matrix says it + tolerates: an explicit null on an optional member (kind "null"), or an + absent key (kind "absent"). A row that becomes null-fragile counts as a + whole, whichever column moved -- a required member made optional while + still throwing on null is one -- and so does a new row that is already + null-fragile, since that is how the defect class keeps arriving. + Everything else -- a fix, stricter type or domain validation, rows or + tests coming and going, harness edits -- is stale. + """ + regressions: list[tuple[str, str, str, str]] = [] + stale: list[str] = [] + old_tests, new_tests = committed["tests"], generated["tests"] + + def where(test: str) -> str: + return f"TEST(JsonPeerInputMatrix, {test})" + + harness = committed["prefix"] != generated["prefix"] + if set(old_tests) == set(new_tests) and committed["tail"] != generated["tail"]: + harness = True + + for test in old_tests: + if test not in new_tests: + stale.append(f"{where(test)}: only in the committed file") + for test, new in new_tests.items(): + old = old_tests.get(test) + if old is None: + stale.append(f"{where(test)}: only in the generated file") + added = new["rows"] + else: + added = {k: r for k, r in new["rows"].items() if k not in old["rows"]} + harness = harness or old["skeleton"] != new["skeleton"] + if old["baseline"] != new["baseline"]: + stale.append(f"{where(test)}: baseline changed") + for name in sorted(set(old["consts"]) | set(new["consts"])): + was, now = old["consts"].get(name), new["consts"].get(name) + if was != now: + stale.append(f"{where(test)}: {name} committed {was}, generated {now}") + for key in old["rows"]: + if key not in new["rows"]: + stale.append(f"{where(test)}: field '{key}' removed") + for key, row in added.items(): + if not _null_fragile(row): + stale.append(f"{where(test)}: field '{key}' added") + for key, prev in old["rows"].items(): + row = new["rows"].get(key) + if row is None: + continue + if prev["present"] != row["present"]: + stale.append(f"{where(test)}: {key} present value changed") + if prev["hint"] != row["hint"]: + stale.append(f"{where(test)}: {key} wrong_hint {prev['hint']}->{row['hint']}") + became_fragile = _null_fragile(row) and not _null_fragile(prev) + for col in COLUMNS: + was, now = prev[col], row[col] + if was == now: + continue + move = f"{where(test)}: {key} {col} {str(was).lower()}->{str(now).lower()}" + if col == "absent" and became_fragile: + # Optional now, but still throwing on null: not a fix. + # When null moved too, its own line names the row. + if prev["null"]: + regressions.append( + ( + test, + key, + f"{where(test)}: {key} became optional but throws " + "on explicit null", + "null", + ) + ) + elif not now: + stale.append(f"{move} (now tolerated)") + elif col == "absent": + regressions.append( + (test, key, f"{move} (member is now required)", "absent") + ) + elif col == "null" and not row["absent"]: + regressions.append( + ( + test, + key, + f"{move} (optional member now throws on explicit null)", + "null", + ) + ) + elif col == "null": + stale.append(f"{move} (required member)") + else: + stale.append(f"{move} (stricter validation)") + for key, row in added.items(): + if _null_fragile(row): + regressions.append( + ( + test, + key, + f"{where(test)}: field '{key}' added as null-fragile " + "(optional member throws on explicit null)", + "null", + ) + ) + + if harness: + stale.append(HARNESS_DIFFERS) + return regressions, stale + + +def regression_block(regressions: list) -> list[str]: + lines = [ + "peer-input matrix REGRESSION: the decoders now throw where the committed " + "matrix says they tolerate the input." + ] + lines += [f" {line}" for _, _, line, _ in regressions] + lines.append( + "Do not regenerate to make this pass; that records the regression as expected behaviour." + ) + kinds = {kind for _, _, _, kind in regressions} + if "null" in kinds: + lines.append("For a null row: make the decoder treat an explicit null as absent.") + if "absent" in kinds: + lines.append( + "For an absent row: the member became required, which breaks every peer that " + "omits it. Check whether the spec now requires it." + ) + lines.append( + "If the stricter behaviour is intended, regenerate with --accept-regression " + "Type.key for each row and say why in the commit message." + ) + return lines + + +def committed_matrix(base: str) -> str | None: + """The matrix as committed at git HEAD, or None when git cannot say.""" + try: + proc = subprocess.run( + ["git", "show", "HEAD:./test/core/json_peer_input_matrix_test.cpp"], + cwd=base, + capture_output=True, + timeout=30, + check=False, + ) + except (OSError, subprocess.TimeoutExpired): + return None + if proc.returncode != 0: + return None + return proc.stdout.decode("utf-8", errors="replace") + + +def guard_baselines(out_text: str | None, head_text: str | None) -> list[str]: + """The matrices a render is guarded against: HEAD's and the output file's. + + Both, because neither alone is enough. Checking only the working copy + lets a row deleted from it come back regressed; checking only HEAD misses + the regression of a fix that is regenerated but not yet committed. Texts + without a TEST are dropped, so the list is empty when there is nothing to + compare with. + """ + baselines: list[str] = [] + for text in (head_text, out_text): + if text is not None and text not in baselines and parse_matrix(text)["tests"]: + baselines.append(text) + return baselines + + +def guard(baselines: list[str], rendered_text: str, accepted: list[str]) -> tuple[int, list[str]]: + """Refuses a render that would record a regression; returns (exit code, lines). + + Regenerating makes the build check pass by definition, so the check alone + cannot stop a regression from being written down as the expected + behaviour. This can: each row that regressed against any of `baselines` + must be named, as ``Type.key`` or as a bare ``Type`` (the TEST name, with + ``::`` spelled ``_``), before it is written. + """ + rendered = parse_matrix(rendered_text) + regressions: list[tuple[str, str, str, str]] = [] + seen: set[tuple[str, str, str]] = set() + for text in baselines: + committed = parse_matrix(text) + if not committed["tests"]: + continue + for test, key, line, kind in classify(committed, rendered)[0]: + if (test, key, kind) not in seen: + seen.add((test, key, kind)) + regressions.append((test, key, line, kind)) + names = {name: name.replace("::", "_") for name in accepted} + lines = [ + f"--accept-regression {name} matches no regressed row" + for name, norm in names.items() + if not any(norm in (test, f"{test}.{key}") for test, key, _, _ in regressions) + ] + blocked = [ + r + for r in regressions + if r[0] not in names.values() and f"{r[0]}.{r[1]}" not in names.values() + ] + if blocked: + return 2, lines + regression_block(blocked) + return (1 if lines else 0), lines + + +def write(text: str, out_path: str) -> None: + with open(out_path, "w", encoding="utf-8") as fh: + fh.write(text) + + # Format on the way out, so regenerating is idempotent against the repo's + # pre-commit hook rather than producing a diff every time. + formatter = shutil.which("clang-format") + if formatter: + subprocess.run([formatter, "-i", out_path], check=False) + + +def main(argv: list[str]) -> int: + ap = argparse.ArgumentParser(description=__doc__) + repo = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + ap.add_argument("--repo", default=repo) + ap.add_argument( + "--out", + default=os.path.join(repo, "test", "core", "json_peer_input_matrix_test.cpp"), + ) + ap.add_argument("--manifest", default=os.path.join(repo, "test", "core", "json_matrix_manifest.json")) + ap.add_argument( + "--accept-regression", + action="append", + default=[], + metavar="Type[.key]", + help="write a row that regressed toward throwing, named by its TEST name " + "(the type, with :: as _) or TEST.key; say why in the commit message", + ) + ap.add_argument( + "--no-baseline", + action="store_true", + help="write even when there is no committed matrix to check for regressions", + ) + args = ap.parse_args(argv) + + text, manifest = render(args.repo) + out_text = None + if os.path.exists(args.out): + with open(args.out, encoding="utf-8") as fh: + out_text = fh.read() + baselines = guard_baselines(out_text, committed_matrix(args.repo)) + if not baselines and not args.no_baseline: + sys.stderr.write( + f"{args.out} is missing or has no TEST(JsonPeerInputMatrix, ...), and git " + "could not supply one from HEAD (git is unavailable, this is not a " + "repository, or no matrix is committed there), so the new rows cannot be " + "checked for regressions. Restore it (git checkout -- " + "test/core/json_peer_input_matrix_test.cpp), or pass --no-baseline.\n" + ) + return 1 + code, lines = guard(baselines, text, args.accept_regression) + for line in lines: + sys.stderr.write(line + "\n") + if code: + return code + + write(text, args.out) + with open(args.manifest, "w", encoding="utf-8") as fh: + # Four-space indent to match the repo's pretty-format-json hook, so + # regenerating does not leave a diff behind. + json.dump(manifest, fh, indent=4, sort_keys=True) + fh.write("\n") + + sys.stderr.write( + f"types={len(manifest['types'])} excluded={len(manifest['excluded'])} " + f"fields={manifest['field_count']} modes={manifest['mode_count']} " + f"cases={manifest['case_count']} gtests={manifest['gtest_count']}\n" + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) diff --git a/scripts/json_census.py b/scripts/json_census.py new file mode 100644 index 0000000..f41ecf4 --- /dev/null +++ b/scripts/json_census.py @@ -0,0 +1,2058 @@ +#!/usr/bin/env python3 +"""Enumerate every deserialization of peer-supplied JSON in the shipped surface. + +The census is regenerated from source on every run. It answers, for each site +where a peer-controlled ``nlohmann::json`` document is read, four questions: + + absent -- is a missing key tolerated? + null -- is a key that is present and explicitly ``null`` tolerated? + wrong_type -- is a key whose value has the wrong JSON type tolerated? + oversized -- is the value's size bounded before it is materialised? + +and one more that decides how much a "no" costs: + + fatality -- does the resulting throw cost one message, one session, or + the process? + +Usage +----- + scripts/json_census.py # census of the working tree + scripts/json_census.py --rev a173cf82 # census of a git revision + scripts/json_census.py --format json # machine-readable rows + scripts/json_census.py --only-flagged # rows that tolerate less than everything + scripts/json_census.py --require FILE:LINE # exit non-zero unless that site is flagged + +``--require`` is the self-test: it asserts that the census rediscovers a site +known to be defective. A census that cannot find a known defect in known +defective code has not been shown to find anything. +""" + +from __future__ import annotations + +import argparse +import csv +import io +import json +import os +import re +import subprocess +import sys +import tarfile +import tempfile +from dataclasses import asdict, dataclass, field +from pathlib import Path +from typing import Iterator, Optional + +# -------------------------------------------------------------------------- +# Tolerance vocabulary +# -------------------------------------------------------------------------- + +OK = "ok" # input in this mode is accepted +THROW = "throw" # input in this mode raises +UNBOUNDED = "unbounded" # no size ceiling is applied before materialising +BOUNDED = "bounded" +NA = "-" + +# Functions that establish a null-tolerant presence test. ``has_json_value`` +# is the helper added in include/mcp/protocol/base.hpp; the census recognises +# it so that it keeps reporting correctly once that helper is in use. +NULL_SAFE_HELPERS = ("has_json_value",) + +# -------------------------------------------------------------------------- +# Exception barriers +# +# A throw is only as expensive as the nearest handler that catches it. These +# tables name the handlers the census knows about; every other function is +# resolved by walking the static call graph. Entries are keyed by the name of +# the function whose body contains the catch. +# -------------------------------------------------------------------------- + +# Catch sites that turn a throw into a JSON-RPC error response and keep going. +PER_MESSAGE_CATCH_MARKERS = ( + "make_error_wire", + "make_error_response", + "send_error", + "McpError", + "g_INVALID_PARAMS", + "g_INVALID_REQUEST", + "g_PARSE_ERROR", + "g_INTERNAL_ERROR", +) + +# Catch sites that tear the session down: the read loop exits, pending requests +# are failed, the transport is closed. +SESSION_FATAL_CATCH_MARKERS = ( + "fail_pending_requests", + "transport->close()", + "closed.store", + "g_CONNECTION_CLOSED", +) + +# Reader loops. A throw that reaches one of these ends the session even when +# the loop catches it, because the loop does not resume. +SESSION_LOOP_FUNCTIONS = ("read_loop", "run_read_loop", "receive_loop", "message_loop") + +FATAL_SESSION = "session" +FATAL_MESSAGE = "message" +FATAL_PROCESS = "process" +FATAL_UNKNOWN = "unknown" + +# -------------------------------------------------------------------------- +# Envelope keys +# +# A JSON-RPC envelope is routed by presence tests, not by get(), so the +# accessor census does not see them. They fail differently too: a presence +# test that misreads an explicit null does not throw, it routes the message +# somewhere wrong -- usually into a silent drop, which is harder to notice +# than a crash. +# +# Null-tolerance is NOT uniform across these keys, and treating them alike is +# its own bug: +# +# "result": null is a legitimate empty result. Presence is the correct +# test; a null check here would reject valid traffic. +# "error": null is not an error. A peer that serialises absent optionals +# as null -- Go without omitempty, a dumped dataclass -- +# sends this routinely, and a presence test misroutes it. +# "id": null is what JSON-RPC prescribes when the id cannot be +# determined; it is not an id. +# "params"/"_meta": null and absent mean the same thing to a peer. +# +# `verdict` is what a *present and null* value does to the routing decision. +# -------------------------------------------------------------------------- + +MISROUTE = "misroute" + +ENVELOPE_KEYS = { + # Keys whose presence test decides where the message goes. + "error": (MISROUTE, "'error': null is not an error but tests as present"), + "id": (MISROUTE, "'id': null is JSON-RPC for an undeterminable id, not an id"), + "result": ("ok", "null is a legitimate empty result; presence is the correct test"), + "method": ("ok", "a null method is malformed whichever way it is tested"), + "jsonrpc": ("ok", "a null jsonrpc is malformed whichever way it is tested"), + # Payload keys. A present null here is stored, not misrouted: the member + # is an optional and ends up engaged holding null. Where the + # guarded block instead calls get() the throw is real, and + # the accessor row for that line already carries it -- recording it twice + # would inflate the count without adding a site. + "params": ("ok", "a present null is stored as an engaged optional, not misrouted"), + "_meta": ("ok", "a present null is stored as an engaged optional, not misrouted"), +} + +# -------------------------------------------------------------------------- +# Source preparation +# -------------------------------------------------------------------------- + +SCAN_ROOTS = ("include", "src") +SOURCE_SUFFIXES = (".hpp", ".cpp", ".h", ".cc") + + +def strip_comments(text: str) -> str: + """Blank out comments while preserving every byte offset and line break. + + Offsets must survive so that reported line numbers match the original file. + String literals are preserved because the JSON keys live in them. + """ + out = list(text) + i = 0 + n = len(text) + while i < n: + ch = text[i] + if ch == '"' or ch == "'": + quote = ch + i += 1 + while i < n: + if text[i] == "\\": + i += 2 + continue + if text[i] == quote: + i += 1 + break + i += 1 + continue + if ch == "/" and i + 1 < n: + if text[i + 1] == "/": + while i < n and text[i] != "\n": + out[i] = " " + i += 1 + continue + if text[i + 1] == "*": + out[i] = out[i + 1] = " " + i += 2 + while i + 1 < n and not (text[i] == "*" and text[i + 1] == "/"): + if text[i] != "\n": + out[i] = " " + i += 1 + if i + 1 < n: + out[i] = out[i + 1] = " " + i += 2 + continue + i += 1 + return "".join(out) + + +def line_index(text: str) -> list[int]: + """Offsets at which each line starts, so an offset maps to a line number.""" + starts = [0] + for m in re.finditer("\n", text): + starts.append(m.end()) + return starts + + +def line_of(starts: list[int], offset: int) -> int: + lo, hi = 0, len(starts) - 1 + while lo < hi: + mid = (lo + hi + 1) // 2 + if starts[mid] <= offset: + lo = mid + else: + hi = mid - 1 + return lo + 1 + + +# -------------------------------------------------------------------------- +# Block structure +# -------------------------------------------------------------------------- + + +@dataclass +class Block: + """One brace-delimited scope, with the text that introduced it.""" + + header: str + start: int # offset just after '{' + end: int # offset of the matching '}' + kind: str # function | if | else | try | catch | loop | other + parent: Optional["Block"] = None + children: list["Block"] = field(default_factory=list) + + +def classify_header(header: str) -> str: + h = header.strip() + if re.search(r"\bcatch\s*\(", h): + return "catch" + if re.search(r"\btry\s*$", h): + return "try" + if re.search(r"\belse\s+if\s*\(", h): + return "if" + if re.search(r"\belse\s*$", h): + return "else" + if re.search(r"^\s*if\s*\(|[});]\s*if\s*\(", h) or re.match(r"^if\s*\(", h): + return "if" + if re.search(r"\b(for|while|switch)\s*\(", h): + return "loop" + if re.search(r"\b(\w+)\s*\([^;]*\)\s*(const\s*)?(noexcept\s*)?(->[^{]*)?$", h) and ( + "=" not in h.split("(")[0] + ): + return "function" + return "other" + + +def parse_blocks(text: str) -> Block: + """Build the brace tree. ``text`` must already have comments blanked.""" + root = Block(header="", start=0, end=len(text), kind="file") + stack = [root] + seg_start = 0 # start of the header accumulating for the next '{' + paren = 0 # a ';' inside parentheses is `for (;;)`, not a statement end + i = 0 + n = len(text) + while i < n: + ch = text[i] + if ch == "(": + paren += 1 + elif ch == ")": + paren = max(0, paren - 1) + if ch == '"' or ch == "'": + quote = ch + i += 1 + while i < n: + if text[i] == "\\": + i += 2 + continue + if text[i] == quote: + i += 1 + break + i += 1 + continue + if ch == "{": + header = text[seg_start:i] + blk = Block( + header=header, + start=i + 1, + end=n, + kind=classify_header(header), + parent=stack[-1], + ) + stack[-1].children.append(blk) + stack.append(blk) + seg_start = i + 1 + paren = 0 + elif ch == "}": + if len(stack) > 1: + blk = stack.pop() + blk.end = i + seg_start = i + 1 + paren = 0 + elif ch == ";" and paren == 0: + seg_start = i + 1 + i += 1 + return root + + +def enclosing_chain(root: Block, offset: int) -> list[Block]: + """Innermost-last list of blocks containing ``offset``.""" + chain: list[Block] = [] + node = root + while True: + for child in node.children: + if child.start <= offset < child.end: + chain.append(child) + node = child + break + else: + return chain + + +# -------------------------------------------------------------------------- +# Guard recognition +# -------------------------------------------------------------------------- + +KEY = r'"((?:[^"\\]|\\.)*)"' + + +def presence_guards(expr: str) -> set[tuple[str, str]]: + """(doc, key) pairs the expression asserts are *present* (may still be null).""" + found = set() + for m in re.finditer(rf"(!\s*)?([\w\->\.\(\)]*?)\.contains\(\s*{KEY}\s*\)", expr): + if m.group(1): # a negated contains() asserts absence, not presence + continue + found.add((normalise_doc(m.group(2)), m.group(3))) + for m in re.finditer(rf"(?:^|[^\w])(\w*::)?find\(\s*{KEY}\s*\)\s*!=", expr): + found.add(("*", m.group(2))) + # A null-safe helper asserts presence as well as non-nullness. Reading it + # only as a null check made every field guarded by one look *required*, + # which would have reported a fixed site as still defective. + for helper in NULL_SAFE_HELPERS: + for m in re.finditer(rf"(?\.\(\)]+)\s*,\s*{KEY}\s*\)", expr): + found.add((normalise_doc(m.group(1)), m.group(2))) + return found + + +def null_safe_guards(expr: str) -> set[tuple[str, str]]: + """(doc, key) pairs the expression asserts are present *and not null*.""" + found = set() + # !doc.at("k").is_null() / !doc["k"].is_null() + for m in re.finditer(rf"!\s*([\w\->\.\(\)]*?)(?:\.at\(\s*{KEY}\s*\)|\[\s*{KEY}\s*\])\s*\.is_null\(\)", expr): + found.add((normalise_doc(m.group(1)), m.group(2) or m.group(3))) + # doc.at("k").is_string() / is_object() / is_number...() -- a positive type + # test implies not-null as a side effect. + for m in re.finditer( + rf"([\w\->\.\(\)]*?)(?:\.at\(\s*{KEY}\s*\)|\[\s*{KEY}\s*\])\s*\.is_(string|object|array|number|number_integer|number_float|number_unsigned|boolean|binary)\(\)", + expr, + ): + found.add((normalise_doc(m.group(1)), m.group(2) or m.group(3))) + for helper in NULL_SAFE_HELPERS: + for m in re.finditer(rf"{helper}\(\s*([\w\->\.\(\)]+)\s*,\s*{KEY}\s*\)", expr): + found.add((normalise_doc(m.group(1)), m.group(2))) + return found + + +KEY_HELPER_SIG = re.compile( + r"\bbool\s+(\w+)\s*\(\s*const\s+(?:nlohmann::)?json\s*&\s*(\w+)\s*\)\s*$" +) + + +def discover_key_helpers(text: str, root: Block) -> dict[str, str]: + """Local predicates that bake the key in, e.g. `has_error_member(message)`. + + ``detail::has_json_value(j, "error")`` names its key, so the guard regexes + see it. A file-local ``has_error_member(message)`` does not, and a census + that only knew the first mechanism would report every site fixed with the + second as still defective -- sending someone to re-fix working code. + Rather than hard-coding the names in use today, a single-return bool + function taking one json document, whose body reduces to exactly one + null-safe test, IS such a helper and is registered as one. + """ + out: dict[str, str] = {} + + def walk(blk: Block) -> None: + m = KEY_HELPER_SIG.search(blk.header.strip()) + if m: + body = text[blk.start : blk.end].strip() + # One return and nothing else: `is_valid_response_envelope` also + # contains null-safe tests but decides several things, so it is not + # a guard for any single key. + if re.fullmatch(r"return\s+[^;]+;", body, re.S): + pairs = null_safe_guards(body) + if len(pairs) == 1: + out[m.group(1)] = next(iter(pairs))[1] + for child in blk.children: + walk(child) + + walk(root) + return out + + +def helper_guards(expr: str, helpers: dict[str, str]) -> set[tuple[str, str]]: + """(doc, key) pairs asserted by a call to a discovered key helper.""" + found = set() + for name, key in helpers.items(): + for m in re.finditer(rf"(?\.]+)\s*\)", expr): + found.add((normalise_doc(m.group(1)), key)) + return found + + +def type_guards(expr: str) -> set[tuple[str, str]]: + """(doc, key) pairs the expression checks the *type* of before use. + + ``is_null()`` is not among them. It establishes that a value is not null + and says nothing about what it is, so reading it as a type check declared + every null-guarded field type-safe -- and an object arriving where a + string-or-integer id was expected still threw. + """ + found = set() + for m in re.finditer( + rf"([\w\->\.\(\)]*?)(?:\.at\(\s*{KEY}\s*\)|\[\s*{KEY}\s*\])\s*\.is_(?!null\b)\w+\(\)", + expr, + ): + found.add((normalise_doc(m.group(1)), m.group(2) or m.group(3))) + return found + + +BOOL_ALIAS = re.compile( + r"(?:const\s+)?bool\s+(\w+)\s*=\s*([^;]+);" +) + + +def bool_aliases(fn_body: str) -> dict[str, str]: + """Locals that stand in for a presence or type test. + + ``const bool has_error = json.contains("error");`` moves the guard out of + the ``if`` header. Without this substitution the census reads the later + ``if (has_error)`` as unguarded and misfiles a null-fragile site as a + required field -- the mistake that let one live site look like a different + defect class. + """ + out: dict[str, str] = {} + for m in BOOL_ALIAS.finditer(fn_body): + name, expr = m.group(1), m.group(2) + if ".contains(" in expr or ".is_" in expr or "has_json_value" in expr: + out[name] = expr + return out + + +def expand_aliases(expr: str, aliases: dict[str, str]) -> str: + """Substitute alias names in a condition with the test they stand for.""" + if not aliases: + return expr + for _ in range(3): + before = expr + for name, replacement in aliases.items(): + expr = re.sub(rf"\b{re.escape(name)}\b", f"({replacement})", expr) + if expr == before: + break + return expr + + +def normalise_doc(raw: str) -> str: + """Reduce a receiver expression to a comparable document name.""" + raw = raw.strip() + raw = re.sub(r"^[\(\!\*&]+", "", raw) + raw = raw.replace("->", ".") + return raw.split(".")[-1] if raw else "*" + + +def doc_matches(a: str, b: str) -> bool: + return a == "*" or b == "*" or a == b + + +# -------------------------------------------------------------------------- +# Peer-document roots +# -------------------------------------------------------------------------- + +JSON_PARAM = re.compile( + r"(?:const\s+)?(?:nlohmann::)?json\s*(?:&|&&|\s)\s*(\w+)\s*(?:,|\))" +) + + +def peer_roots_for(block: Block, text: str) -> set[str]: + """Names in scope that hold a document a peer controls. + + Roots are function parameters of ``nlohmann::json`` type; the set then + propagates through local bindings initialised from an existing root. A + function parameter typed ``nlohmann::json`` is peer-controlled because every + such parameter in this codebase is reached from ``json::parse`` of bytes + off the wire; the census states that assumption rather than proving it, and + ``--list-roots`` prints the set so it can be audited. + """ + roots: set[str] = set() + fn = block + while fn is not None and fn.kind != "function": + fn = fn.parent + if fn is None: + return roots + for m in JSON_PARAM.finditer(fn.header): + roots.add(m.group(1)) + if not roots: + return roots + body = text[fn.start : fn.end] + # Propagate through local bindings: `auto x = ` + for _ in range(4): # fixed point; depth 4 is past anything in this tree + before = len(roots) + pattern = re.compile( + r"(?:const\s+)?(?:auto|nlohmann::json|json)\s*(?:&|&&)?\s*(\w+)\s*=\s*([^;]+);" + ) + for m in pattern.finditer(body): + name, init = m.group(1), m.group(2) + if name in roots: + continue + if any(re.search(rf"\b{re.escape(r)}\b", init) for r in roots): + roots.add(name) + if len(roots) == before: + break + return roots + + +# -------------------------------------------------------------------------- +# Access sites +# -------------------------------------------------------------------------- + +ACCESS_PATTERNS = [ + # doc.at("k").get() / .get_to(...) / .get_ref<...> + ( + "at_get", + re.compile( + rf"([\w\->\.]+?)\.at\(\s*{KEY}\s*\)\s*\.\s*(get|get_to|get_ref|get_ptr)\b" + ), + ), + # doc["k"].get() + ( + "sub_get", + re.compile(rf"([\w\->\.]+?)\[\s*{KEY}\s*\]\s*\.\s*(get|get_to|get_ref|get_ptr)\b"), + ), + # doc.at("k") used as a value (assignment / argument), no explicit get + ("at_bare", re.compile(rf"([\w\->\.]+?)\.at\(\s*{KEY}\s*\)")), + # doc.value("k", default) + ("value_default", re.compile(rf"([\w\->\.]+?)\.value\(\s*{KEY}\s*,")), + # doc.get() on the whole document + ("doc_get", re.compile(r"([\w\->\.]+?)\.get(?:_to)?\s*<")), +] + +TYPE_ARG = re.compile(r"\.get(?:_to|_ref|_ptr)?\s*<\s*([^>]+(?:<[^>]*>)?[^>]*)\s*>") + +SIZE_LIMIT_TOKENS = ( + ".size() >", + ".size() >=", + ".length() >", + "max_size", + "max_message", + "max_body", + "size_limit", + "MAX_", + "g_MAX", +) + + +@dataclass +class Row: + file: str + line: int + owner: str # deserialised type, or the enclosing function + field_name: str + site_kind: str + absent: str + null: str + wrong_type: str + oversized: str + guard: str + fatality: str + function: str + snippet: str + target_types: tuple[str, ...] = () + fatality_basis: str = "lexical" + target_hint: str = "" # the get the site names, verbatim + member: str = "" # the C++ member the key is read into + constrained: str = NA # the field has a value domain, not just a type + + @property + def flagged(self) -> bool: + """The site rejects, or misroutes, input a careful peer may legitimately send.""" + return THROW in (self.absent, self.null, self.wrong_type) or self.null == MISROUTE + + @property + def defect_class(self) -> str: + """The shape of the site's intolerance, independent of its line number.""" + if self.null == MISROUTE: + return "misroute" + if self.absent == OK and self.null == THROW: + # Absent is handled, present-and-null is not. This is the class + # that has outrun prediction: a peer that serialises absent + # optionals as explicit null reaches it without trying. + return "null-fragile" + if self.absent == THROW: + return "required" + if self.wrong_type == THROW: + return "untyped" + return "tolerant" + + +def statement_end(text: str, offset: int, limit: int) -> int: + """Offset of the ';' that ends the statement containing ``offset``. + + A fixed character window is not a statement. Searching 200 characters + ahead for a `get` picked up the *next* line's type whenever the current + line read through `get_to(...)`, which names no type -- so `roots` was + typed from the `_meta` line below it, and the verdict came out inverted. + """ + depth = 0 + i = offset + while i < limit: + ch = text[i] + if ch == "(": + depth += 1 + elif ch == ")": + depth = max(0, depth - 1) + elif ch == ";" and depth == 0: + return i + elif ch in "{}" and depth == 0: + return i + i += 1 + return limit + + +def statement_prefix(text: str, offset: int, block: Block) -> str: + """Text of the current statement up to ``offset``, for same-expression guards.""" + start = block.start + for sep in (";", "{", "}"): + idx = text.rfind(sep, block.start, offset) + if idx != -1: + start = max(start, idx + 1) + return text[start:offset] + + +def enclosing_function(chain: list[Block]) -> tuple[str, Optional[Block]]: + for blk in reversed(chain): + if blk.kind == "function": + m = re.search(r"(\w+)\s*\([^()]*\)\s*(?:const\s*)?(?:noexcept\s*)?(?:->[^{]*)?$", blk.header.strip()) + if m: + return m.group(1), blk + return "?", blk + return "", None + + +FROM_JSON_SIG = re.compile( + r"from_json\s*\(\s*const\s+(?:nlohmann::)?json\s*&\s*(\w+)\s*,\s*([\w:]+(?:<[^>]*>)?)\s*&" +) + + +def deserialised_type(chain: list[Block]) -> Optional[str]: + for blk in reversed(chain): + if blk.kind == "function": + m = FROM_JSON_SIG.search(blk.header) + if m: + return m.group(2) + return None + + +def size_bounded(fn_body: str) -> bool: + return any(tok in fn_body for tok in SIZE_LIMIT_TOKENS) + + +# -------------------------------------------------------------------------- +# Fatality resolution +# -------------------------------------------------------------------------- + + +@dataclass +class FunctionInfo: + name: str + file: str + block: Block + body: str + calls: set[str] + + +def unit_key(rel_path: str) -> str: + """Header and implementation of one component count as one unit. + + `client.hpp` declares what `client.cpp` defines and calls; resolving the + call graph per-file would sever that pair. + """ + return os.path.splitext(os.path.basename(rel_path))[0] + + +def catch_disposition(catch_body: str) -> str: + if any(tok in catch_body for tok in SESSION_FATAL_CATCH_MARKERS): + return FATAL_SESSION + if any(tok in catch_body for tok in PER_MESSAGE_CATCH_MARKERS): + return FATAL_MESSAGE + if re.search(r"\bthrow\s*;", catch_body): + return "rethrow" + return FATAL_MESSAGE # a catch that neither tears down nor rethrows contains the throw + + +def local_barrier(text: str, chain: list[Block]) -> Optional[str]: + """Disposition of the innermost try/catch lexically containing the site. + + The structural rule outranks anything the catch body says. A ``try`` that + *encloses* the read loop cannot resume it: control leaves the loop, no + further message is read, and the session is over whether or not the handler + says so. An empty catch around ``for (;;) { ... }`` is session-fatal. A + ``try`` *inside* the loop is a per-message barrier: the next iteration runs. + """ + for blk in reversed(chain): + if blk.kind != "try": + continue + parent = blk.parent + if parent is None: + continue + dispositions = [] + for sib in parent.children: + if sib.kind == "catch" and sib.start > blk.end: + dispositions.append(catch_disposition(text[sib.start : sib.end])) + if not dispositions or all(d == "rethrow" for d in dispositions): + continue + # Is there a loop between this try and the site? If so the catch is + # outside the loop and unwinding ends it. + idx = chain.index(blk) + if any(b.kind == "loop" for b in chain[idx + 1 :]): + return FATAL_SESSION + if FATAL_SESSION in dispositions: + return FATAL_SESSION + return FATAL_MESSAGE + return None + + +def resolve_fatality( + text: str, + chain: list[Block], + fn_name: str, + index: dict[str, list[FunctionInfo]], + current_file: str, + depth: int = 0, + seen: Optional[set[str]] = None, +) -> str: + """Nearest handler on any static path out of the site.""" + local = local_barrier(text, chain) + if local is not None: + return local + if fn_name in SESSION_LOOP_FUNCTIONS: + return FATAL_SESSION + seen = seen or set() + if fn_name in seen or depth > 6: + return FATAL_UNKNOWN + seen.add(fn_name) + + # Prefer callers in the same translation unit. `dispatch_response` exists + # in both the client and the server; resolving by name alone would let the + # client's read loop decide the server's blast radius. + unit = unit_key(current_file) + same_file = { + name: [i for i in infos if unit_key(i.file) == unit] + for name, infos in index.items() + } + scope = {n: i for n, i in same_file.items() if i} + if not any( + re.search(rf"\b{re.escape(fn_name)}\s*\(", i.body) for l in scope.values() for i in l + ): + scope = index + + outcomes = set() + for caller_name, infos in scope.items(): + for info in infos: + for m in re.finditer(rf"\b{re.escape(fn_name)}\s*\(", info.body): + call_off = info.block.start + m.start() + call_chain = enclosing_chain(info.block, call_off) + sub = local_barrier(info.body_text, [info.block] + call_chain) # type: ignore[attr-defined] + if sub is not None: + outcomes.add(sub) + elif caller_name in SESSION_LOOP_FUNCTIONS: + outcomes.add(FATAL_SESSION) + else: + outcomes.add( + resolve_fatality( + info.body_text, # type: ignore[attr-defined] + [info.block] + call_chain, + caller_name, + index, + info.file, + depth + 1, + set(seen), + ) + ) + if not outcomes: + return FATAL_UNKNOWN + # Report the worst barrier actually found. A chain that also has an + # unresolved branch is still reported at its known severity rather than + # collapsing to "unknown", because one unreachable caller must not erase + # the evidence from the reachable ones. `--incomplete-chains` lists the + # rows where some branch stayed unresolved. + for level in (FATAL_PROCESS, FATAL_SESSION, FATAL_MESSAGE): + if level in outcomes: + return level + return FATAL_UNKNOWN + + +# -------------------------------------------------------------------------- +# Scanning +# -------------------------------------------------------------------------- + +ENUM_MACRO = re.compile( + r"NLOHMANN_JSON_SERIALIZE_ENUM\s*\(\s*([\w:]+)\s*,\s*\{(.*?)\}\s*\)", + re.DOTALL, +) + +# A validator called on a member turns a field into a value constraint: the +# field has a domain, not just a JSON type, so an oversized or arbitrary value +# of the right type is still rejected. Without this the matrix reads such a +# rejection as a size limit that does not exist. +VALIDATOR_CALL = re.compile(r"\b(?:detail::)?(validate_\w+)\s*\(\s*([\w\.\->]+)") + +MACRO_DEFINE = re.compile( + r"NLOHMANN_DEFINE_TYPE_(NON_INTRUSIVE|INTRUSIVE)(_WITH_DEFAULT)?\s*\(\s*([\w:]+)\s*,\s*([^)]*)\)" +) + + +# A handler stored in a map and invoked through `iter->second(...)` is an +# indirect call. No name-based call graph crosses it -- which is why a +# progress callback could sit on the read loop's blast radius while every +# static reading of it said "unknown". These two patterns recover the edge: +# where a lambda is stored, and where the container that holds it is called. +REGISTRY_LOOKUP = re.compile( + r"\w*(?:handler|callback)s?\s*\.(?:find|at|contains)\s*\(" +) + + +def indirect_dispatch_radius( + text: str, + root: Block, + index: dict[str, list[FunctionInfo]], + rel: str, +) -> str: + """Blast radius of the worst function that dispatches to a stored handler. + + Discovery runs from the dispatch side, not the registration side. A + function that looks a key up in a handler registry goes on to call what it + found -- whether as ``iter->second(...)`` or, as here, by copying the + callable out under a lock and calling it afterwards. Registration may be + several hops away (``on_progress`` -> ``on_notification`` -> + ``handlers[m] = cb``) and need not be followed: every handler in the + registry shares the radius of the function that dispatches to it. + """ + radius = FATAL_UNKNOWN + seen: set[int] = set() + for lookup in REGISTRY_LOOKUP.finditer(text): + chain = enclosing_chain(root, lookup.start()) + if not chain: + continue + fn = next((b for b in reversed(chain) if b.kind == "function"), None) + if fn is None or fn.start in seen: + continue + seen.add(fn.start) + fn_chain = enclosing_chain(root, fn.start) + fn_name, _ = enclosing_function(fn_chain) + radius = worst(radius, resolve_fatality(text, fn_chain, fn_name, index, rel)) + return radius + + +JSON_LAMBDA = re.compile(r"\[[^\]\[]*\]\s*\([^)]*(?:const\s+)?(?:nlohmann::)?json\s*&") + + +def peer_handler_regions(text: str, root: Block, radius: str) -> list[tuple[int, int, str]]: + """Offset ranges of lambdas that receive a peer document.""" + if radius == FATAL_UNKNOWN: + return [] + regions: list[tuple[int, int, str]] = [] + for m in JSON_LAMBDA.finditer(text): + blk = lambda_block_after(root, m.start()) + if blk is not None: + regions.append((blk.start, blk.end, radius)) + return regions + + +def fatality_at( + offset: int, lexical: str, regions: list[tuple[int, int, str]] +) -> str: + """Upgrade an unresolved site that sits inside a peer-invoked handler.""" + if lexical != FATAL_UNKNOWN: + return lexical + for start, end, radius in regions: + if start <= offset < end: + return radius + return lexical + + +def lambda_block_after(root: Block, offset: int) -> Optional[Block]: + """The block that opens first after ``offset`` -- a stored lambda's body.""" + best: Optional[Block] = None + + def walk(blk: Block) -> None: + nonlocal best + if blk.start > offset and (best is None or blk.start < best.start): + best = blk + for child in blk.children: + walk(child) + + walk(root) + return best + + +def scan_file( + path: str, rel: str, index: dict[str, list[FunctionInfo]] +) -> tuple[list[Row], dict[str, list[tuple[str, str]]]]: + raw = Path(path).read_text(encoding="utf-8", errors="replace") + text = strip_comments(raw) + starts = line_index(text) + root = parse_blocks(text) + structs = parse_structs(text, root) + helpers = discover_key_helpers(text, root) + indirect_regions = peer_handler_regions( + text, root, indirect_dispatch_radius(text, root, index, rel) + ) + rows: list[Row] = [] + + # -- macro-generated serialisers --------------------------------------- + for m in MACRO_DEFINE.finditer(text): + with_default = bool(m.group(2)) + type_name = m.group(3) + fields = [f.strip() for f in m.group(4).split(",") if f.strip()] + line = line_of(starts, m.start()) + for fname in fields: + rows.append( + Row( + file=rel, + line=line, + owner=type_name, + field_name=fname, + site_kind="macro_with_default" if with_default else "macro_required", + absent=OK if with_default else THROW, + # the macro emits at(...).get_to(...); a present null throws + # for every member type but nlohmann::json and std::optional + null=THROW, + wrong_type=THROW, + oversized=UNBOUNDED, + guard="none (macro)", + fatality=FATAL_UNKNOWN, + function=f"", + snippet=m.group(0)[:120], + target_types=named_types( + next((t for t, n in structs.get(type_name, []) if n == fname), "") + ), + # NLOHMANN_DEFINE_TYPE keys each member by its own name. + member=fname, + ) + ) + + # -- explicit accessor sites ------------------------------------------- + claimed: set[tuple[int, str]] = set() + for pattern_kind, pattern in ACCESS_PATTERNS: + for m in pattern.finditer(text): + site_kind = pattern_kind + off = m.start() + doc_raw = m.group(1) + doc = normalise_doc(doc_raw) + key = m.group(2) if pattern.groups >= 2 and site_kind != "doc_get" else None + + chain = enclosing_chain(root, off) + if not chain: + continue + fn_name, fn_block = enclosing_function(chain) + if fn_block is None: + continue + roots = peer_roots_for(chain[-1], text) + if doc not in roots: + continue + + # de-duplicate: at_bare also matches what at_get already claimed + marker = (off, doc) + if site_kind == "at_bare" and any( + c[0] == off for c in claimed + ): + continue + claimed.add(marker) + + line = line_of(starts, off) + stmt = statement_prefix(text, off, chain[-1]) + + # Guards in force: enclosing if-conditions, plus the current + # statement's own prefix (ternaries and && chains). + aliases = bool_aliases(text[fn_block.start : fn_block.end]) + present: set[tuple[str, str]] = set() + nonnull: set[tuple[str, str]] = set() + typed: set[tuple[str, str]] = set() + rejected: set[tuple[str, str]] = set() + for blk in chain: + if blk.kind in ("if", "loop"): + header = expand_aliases(blk.header, aliases) + present |= presence_guards(header) + hg = helper_guards(header, helpers) + present |= hg + body = text[blk.start : blk.end] + if terminates(body) and re.search(r"!\s*[\w\->\.]+?(?:\.at\(|\[)", header): + # `if (contains(k) && !at(k).is_object()) { throw; }` + # proves the shape for the code that follows, but what + # the peer sees is a rejection: sending null here + # throws. Counting it as a tolerance inverted the + # verdict on every validated field. + rejected |= type_guards(header) + else: + nonnull |= null_safe_guards(header) + nonnull |= hg + typed |= type_guards(header) + stmt_x = expand_aliases(stmt, aliases) + present |= presence_guards(stmt_x) + nonnull |= null_safe_guards(stmt_x) + present |= helper_guards(stmt_x, helpers) + nonnull |= helper_guards(stmt_x, helpers) + typed |= type_guards(stmt_x) + nonnull -= rejected + typed -= rejected + present |= early_return_guards(text, chain, off, aliases) + nonnull |= early_return_null_guards(text, chain, off, helpers) + + if key is None: + # A whole-document get(): the document is handed to T's own + # from_json, whose fields this census enumerates separately. + # The row records the delegation and whether the document + # itself was shape-checked first. + key = "" + guarded_present = True + guarded_null = any( + re.search(r"is_null\(\)|is_object\(\)|is_array\(\)", b.header) + for b in chain + if b.kind == "if" + ) or "is_null()" in stmt or "is_object()" in stmt + guarded_type = guarded_null + site_kind = "delegate" + else: + guarded_present = any(doc_matches(d, doc) and k == key for d, k in present) + guarded_null = any(doc_matches(d, doc) and k == key for d, k in nonnull) + guarded_type = any(doc_matches(d, doc) and k == key for d, k in typed) + + stmt_end = statement_end(text, off, min(len(text), off + 400)) + target_type = TYPE_ARG.search(text[off:stmt_end]) + tt = target_type.group(1).strip() if target_type else "" + # Only a bare nlohmann::json absorbs anything. `std::map` does not: constructing the map from a null + # throws, so a substring test here declared a whole class of + # fields null-tolerant that are not. + bare_tt = re.sub(r"^(?:const\s+)?std::optional\s*<(.+)>$", r"\1", tt.strip()).strip() + json_typed = bare_tt in ("nlohmann::json", "json") or site_kind == "at_bare" + + if site_kind == "delegate": + absent = NA + null = OK if guarded_null or json_typed else THROW + wrong = OK if guarded_type or json_typed else THROW + guard_desc = ( + f"delegates to {tt or '?'}::from_json" + + ("" if guarded_null else "; document not shape-checked") + ) + elif site_kind == "value_default": + absent = OK + null = OK if json_typed else THROW + wrong = THROW + guard_desc = "value(key, default)" + else: + absent = OK if guarded_present else THROW + if json_typed: + null = OK + elif guarded_null: + null = OK + else: + null = THROW + wrong = OK if (guarded_type or json_typed) else THROW + if guarded_null: + guard_desc = "contains + null/type check" + elif guarded_present: + guard_desc = "contains only" + else: + guard_desc = "none (at)" + + fn_body = text[fn_block.start : fn_block.end] + member = member_target(text, off, stmt_end, stmt) + constrained = NA + if member: + for vm in VALIDATOR_CALL.finditer(fn_body): + arg = vm.group(2).replace("->", ".") + if arg.split(".")[-1] == member: + constrained = vm.group(1) + break + rows.append( + Row( + file=rel, + line=line, + owner=deserialised_type(chain) or f"({fn_name})", + field_name=key, + site_kind=site_kind, + absent=absent, + null=null, + wrong_type=wrong, + oversized=BOUNDED if size_bounded(fn_body) else UNBOUNDED, + guard=guard_desc, + fatality=fatality_at( + off, resolve_fatality(text, chain, fn_name, index, rel), indirect_regions + ), + function=fn_name, + snippet=raw.splitlines()[line - 1].strip()[:140] + if line - 1 < len(raw.splitlines()) + else "", + target_types=named_types(tt), + target_hint=tt, + member=member, + constrained=constrained, + ) + ) + + rows.extend( + scan_envelope_predicates( + text, raw, starts, root, index, rel, indirect_regions, helpers + ) + ) + return rows, structs + + +CONTAINS_CALL = re.compile(rf"([\w\->\.]+?)\.contains\(\s*{KEY}\s*\)") + + +def scan_envelope_predicates( + text: str, + raw: str, + starts: list[int], + root: Block, + index: dict[str, list[FunctionInfo]], + rel: str, + indirect_regions: list[tuple[int, int, str]], + helpers: dict[str, str], +) -> list[Row]: + """Presence tests on envelope keys, and what an explicit null does to them.""" + lines = raw.splitlines() + out: list[Row] = [] + for m in CONTAINS_CALL.finditer(text): + key = m.group(2) + if key not in ENVELOPE_KEYS: + continue + off = m.start() + chain = enclosing_chain(root, off) + if not chain: + continue + fn_name, fn_block = enclosing_function(chain) + if fn_block is None: + continue + doc = normalise_doc(m.group(1)) + if doc not in peer_roots_for(chain[-1], text): + continue + + verdict, reason = ENVELOPE_KEYS[key] + stmt = statement_prefix(text, m.end(), chain[-1]) + window = text[max(fn_block.start, off - 200) : off + 200] + aliases = bool_aliases(text[fn_block.start : fn_block.end]) + expanded = expand_aliases(window, aliases) + checked = any( + doc_matches(d, doc) and k == key + for d, k in (null_safe_guards(expanded) | helper_guards(expanded, helpers)) + ) + if checked: + verdict, reason = "ok", "presence test is paired with a null check" + + line = line_of(starts, off) + out.append( + Row( + file=rel, + line=line, + owner=f"({fn_name})", + field_name=key, + site_kind="envelope_presence", + absent=OK, + null=OK if verdict == "ok" else MISROUTE, + wrong_type=NA, + oversized=NA, + guard=reason, + # A misroute does not throw, so throw-radius does not describe + # it: the message is routed wrong and, on a correlation path, + # silently dropped -- the pending request then hangs to its + # timeout rather than failing loudly. + fatality=MISROUTE + if verdict == MISROUTE + else fatality_at( + off, + resolve_fatality(text, chain, fn_name, index, rel), + indirect_regions, + ), + function=fn_name, + snippet=lines[line - 1].strip()[:140] if line - 1 < len(lines) else "", + ) + ) + return out + + +def terminates(body: str) -> bool: + """Does this block leave the enclosing function or loop iteration?""" + return bool(re.search(r"\b(return|co_return|throw|continue|break)\b", body)) + + +GET_TO_TARGET = re.compile(r"\.get_to\s*\(\s*[\w]+\s*(?:\.|->)\s*(\w+)\s*\)") +ASSIGN_TARGET = re.compile(r"(?:\.|->)\s*(\w+)\s*=\s*$") + + +def member_target(text: str, offset: int, end: int, stmt: str) -> str: + """The C++ member a key is read into, so the field can be typed. + + The wire key and the member name usually match, but not always -- `Icon` + reads its `source` member from `"src"`. Guessing from the member name + produced a baseline document that was missing a required key, which made + every mode of every field of that type throw for a reason unrelated to + what the case was testing. + """ + m = GET_TO_TARGET.search(text[offset:end]) + if m: + return m.group(1) + m = ASSIGN_TARGET.search(stmt.rstrip()) + if m: + return m.group(1) + return "" + + +def early_return_guards( + text: str, chain: list[Block], offset: int, aliases: dict[str, str] +) -> set[tuple[str, str]]: + """Keys proven present by a preceding `if (!doc.contains(k)) return;`.""" + found = set() + for blk in chain: + for sib in blk.children: + if sib.kind != "if" or sib.end >= offset: + continue + body = text[sib.start : sib.end] + if not re.search(r"\b(return|co_return|throw|continue|break)\b", body): + continue + header = expand_aliases(sib.header, aliases) + for m in re.finditer(rf"!\s*([\w\->\.]+?)\.contains\(\s*{KEY}\s*\)", header): + found.add((normalise_doc(m.group(1)), m.group(2))) + # `if (!(a && doc.contains(k)))` and `if (!a || !doc.contains(k))` + for m in re.finditer(rf"\|\|\s*!\s*([\w\->\.]+?)\.contains\(\s*{KEY}\s*\)", header): + found.add((normalise_doc(m.group(1)), m.group(2))) + return found + + +def early_return_null_guards( + text: str, chain: list[Block], offset: int, helpers: dict[str, str] +) -> set[tuple[str, str]]: + """Keys proven non-null by a preceding early-return type/null test.""" + found = set() + for blk in chain: + for sib in blk.children: + if sib.kind != "if" or sib.end >= offset: + continue + body = text[sib.start : sib.end] + if not re.search(r"\b(return|co_return|throw|continue|break)\b", body): + continue + for m in re.finditer( + rf"!\s*([\w\->\.]+?)(?:\.at\(\s*{KEY}\s*\)|\[\s*{KEY}\s*\])\s*\.is_\w+\(\)", + sib.header, + ): + found.add((normalise_doc(m.group(1)), m.group(2) or m.group(3))) + for m in re.finditer( + rf"([\w\->\.]+?)(?:\.at\(\s*{KEY}\s*\)|\[\s*{KEY}\s*\])\s*\.is_null\(\)", + sib.header, + ): + found.add((normalise_doc(m.group(1)), m.group(2) or m.group(3))) + # `if (!has_error_member(m)) { return; }` proves, for everything + # after it, that the key is present and not null. + for name, key in helpers.items(): + for m in re.finditer( + rf"!\s*{re.escape(name)}\s*\(\s*([\w\->\.]+)\s*\)", sib.header + ): + found.add((normalise_doc(m.group(1)), key)) + return found + + +def build_index(files: list[tuple[str, str]]) -> dict[str, list[FunctionInfo]]: + """Name -> function bodies, for the call-graph walk used by fatality.""" + index: dict[str, list[FunctionInfo]] = {} + for path, rel in files: + raw = Path(path).read_text(encoding="utf-8", errors="replace") + text = strip_comments(raw) + root = parse_blocks(text) + + def walk(blk: Block) -> None: + if blk.kind == "function": + m = re.search( + r"(\w+)\s*\([^()]*\)\s*(?:const\s*)?(?:noexcept\s*)?(?:->[^{]*)?$", + blk.header.strip(), + ) + if m: + info = FunctionInfo( + name=m.group(1), + file=rel, + block=blk, + body=text[blk.start : blk.end], + calls=set(), + ) + info.body_text = text # type: ignore[attr-defined] + index.setdefault(m.group(1), []).append(info) + for child in blk.children: + walk(child) + + walk(root) + return index + + +# -------------------------------------------------------------------------- +# Struct members and the type graph +# +# nlohmann reaches a type's ``from_json`` through ``get()``, which no +# name-based call graph can see: there is no call to `from_json` in the text. +# Without this edge every protocol serialiser reports an unknown blast radius, +# which is the same as reporting nothing. The type graph supplies the edge. +# -------------------------------------------------------------------------- + +STRUCT_DECL = re.compile(r"\b(?:struct|class)\s+([A-Z]\w*)\s*(?::([^{]*))?\{") +BASE_NAME = re.compile(r"\b(?:public|protected|private)?\s*([A-Z]\w*)") +MEMBER_DECL = re.compile( + r"^\s*((?:const\s+)?[\w:]+(?:\s*<[^;]*>)?)\s+(\w+)\s*(?:=\s*[^;]+)?;\s*$", + re.MULTILINE, +) + + +def parse_bases(text: str, root: Block) -> dict[str, list[str]]: + """Type name -> its base classes. + + A derived serialiser delegates to its base -- EnumSchema's from_json calls + the PrimitiveSchemaDefinition one -- so the base's required fields are + required of the derived type too. Without this edge the matrix built a + baseline missing an inherited key and every case for that type threw on + the inherited field rather than on the mode under test. + """ + out: dict[str, list[str]] = {} + + def walk(blk: Block) -> None: + m = STRUCT_DECL.search(blk.header + "{") + if m and m.group(2): + out.setdefault(m.group(1), BASE_NAME.findall(m.group(2))) + for child in blk.children: + walk(child) + + walk(root) + return out + + +def parse_structs(text: str, root: Block) -> dict[str, list[tuple[str, str]]]: + """Type name -> [(member type, member name)] for every struct in the file.""" + out: dict[str, list[tuple[str, str]]] = {} + + def walk(blk: Block) -> None: + m = STRUCT_DECL.search(blk.header + "{") + if m: + name = m.group(1) + body = text[blk.start : blk.end] + # Only direct members: drop anything inside a nested brace. + flat = strip_nested(body) + members = [ + (t.strip(), n) + for t, n in MEMBER_DECL.findall(flat) + if t.strip() not in ("return", "co_return") + ] + if members or name not in out: + out[name] = members + for child in blk.children: + walk(child) + + walk(root) + return out + + +def strip_nested(body: str) -> str: + """Blank out everything inside nested braces, preserving line structure.""" + out = list(body) + depth = 0 + for i, ch in enumerate(body): + if ch == "{": + depth += 1 + out[i] = " " + elif ch == "}": + depth = max(0, depth - 1) + out[i] = " " + elif depth > 0 and ch != "\n": + out[i] = " " + return "".join(out) + + +IDENT = re.compile(r"\b([A-Z]\w*(?:::\w+)*)\b") + +CONTAINER_NOISE = {"T", "Result", "Params", "Args"} + + +def named_types(type_expr: str) -> tuple[str, ...]: + """Capitalised type names mentioned in a template argument. + + ``std::vector`` yields ``RelatedTaskMetadata``; + ``std::map`` yields nothing. + """ + if not type_expr: + return () + return tuple( + t + for t in IDENT.findall(type_expr) + if t not in CONTAINER_NOISE and not t.startswith("std::") + ) + +FATAL_ORDER = [FATAL_PROCESS, FATAL_SESSION, FATAL_MESSAGE, FATAL_UNKNOWN] + + +def worst(a: str, b: str) -> str: + for f in FATAL_ORDER: + if a == f or b == f: + return f + return FATAL_UNKNOWN + + +def propagate_type_fatality( + rows: list[Row], + delegations: dict[str, set[str]], + deserialisable: set[str], +) -> None: + """Give every protocol row the blast radius of the worst entry that reaches it. + + An entry is a site outside the protocol headers that hands a document to + ``get()``; its blast radius is already resolved lexically. That radius + then flows to T and to every type T's serialiser delegates to. + """ + seeds: dict[str, str] = {} + for r in rows: + if r.fatality == FATAL_UNKNOWN: + continue + for t in r.target_types: + if t in deserialisable: + seeds[t] = worst(seeds.get(t, FATAL_UNKNOWN), r.fatality) + + resolved = dict(seeds) + frontier = list(seeds.items()) + guard = 0 + while frontier and guard < 10000: + guard += 1 + t, f = frontier.pop() + for nxt in delegations.get(t, ()): # noqa: B007 + new = worst(resolved.get(nxt, FATAL_UNKNOWN), f) + if resolved.get(nxt) != new: + resolved[nxt] = new + frontier.append((nxt, new)) + + for r in rows: + if r.fatality == FATAL_UNKNOWN and r.owner in resolved: + r.fatality = resolved[r.owner] + r.fatality_basis = "reached" + + +def strict_serialisers(files: list[tuple[str, str]]) -> set[str]: + """Types whose from_json raises rather than accepting whatever it is given.""" + strict: set[str] = set() + for path, _rel in files: + text = strip_comments(Path(path).read_text(encoding="utf-8", errors="replace")) + root = parse_blocks(text) + + def walk(blk: Block) -> None: + m = FROM_JSON_SIG.search(blk.header) + if m and re.search(r"\bthrow\b", text[blk.start : blk.end]): + strict.add(m.group(2).split("::")[-1]) + for child in blk.children: + walk(child) + + walk(root) + return strict + + +def enum_table(text: str) -> dict[str, str]: + """Enum type -> its first wire string, which is also its fallback value.""" + out: dict[str, str] = {} + for m in ENUM_MACRO.finditer(text): + first = re.search(r'"([^"]+)"', m.group(2)) + if first: + out[m.group(1).split("::")[-1]] = first.group(1) + return out + + +def fill_target_hints( + rows: list[Row], structs: dict[str, list[tuple[str, str]]] +) -> None: + """Type a site that reads through ``get_to(x.member)`` from the member. + + ``at("data").get_to(params.data)`` names no type at the call site, so the + site alone cannot say whether a null is tolerated. The member's + declaration can: ``nlohmann::json data`` absorbs a null, ``TaskStatus + status`` coerces one, ``std::string`` rejects it. + """ + for r in rows: + if r.target_hint or not r.member or not r.owner: + continue + for ctype, mname in structs.get(r.owner, []): + if mname == r.member: + r.target_hint = ctype + r.target_types = r.target_types or named_types(ctype) + break + + +def refine_targets( + rows: list[Row], + enums: dict[str, str], + required_types: set[str], + known: set[str], + strict: set[str], +) -> None: + """Correct verdicts that depend on what the field deserialises *into*. + + Two cases read as a throw from the call site but are not: + + * An enum built with NLOHMANN_JSON_SERIALIZE_ENUM does not throw on an + unrecognised value -- it falls back to the first enumerator. Null and + wrong-typed input are accepted and silently coerced, which is a quieter + failure than a throw and worth naming as its own class. + * A type whose serialiser reads only optional fields accepts a null + document: `contains()` on a null returns false for every key, so every + field is skipped and an empty value is produced. + """ + for r in rows: + if r.site_kind in ("delegate", "envelope_presence", "macro_required"): + continue + # The refinement is about what the field deserialises into as a whole. + # `std::vector` is not a Role: constructing the vector from a + # null throws before any enum coercion can happen, so matching on a + # template argument inverted the verdict for every container of an + # enum or of an all-optional struct. + whole = re.sub(r"^(?:const\s+)?std::optional\s*<(.+)>$", r"\1", r.target_hint.strip()).strip() + target = whole.split("::")[-1] if whole else None + if target is None: + continue + if whole not in ("nlohmann::json", "json") and target not in enums and target not in known: + continue + if whole in ("nlohmann::json", "json"): + r.null = OK + r.wrong_type = OK + r.guard += "; the member is a raw json value and absorbs anything" + elif target in enums: + r.null = OK + r.wrong_type = OK + r.guard += "; enum coerces an unrecognised value to the first enumerator" + elif target not in required_types and target not in strict: + r.null = OK + r.wrong_type = OK + r.guard += "; target has no required field, so a null parses as an empty value" + + +def collect_files(base: str) -> list[tuple[str, str]]: + files = [] + for rootdir in SCAN_ROOTS: + top = os.path.join(base, rootdir) + if not os.path.isdir(top): + continue + for dirpath, _dirs, names in os.walk(top): + for name in sorted(names): + if name.endswith(SOURCE_SUFFIXES): + full = os.path.join(dirpath, name) + files.append((full, os.path.relpath(full, base))) + return sorted(files, key=lambda p: p[1]) + + +def export_revision(repo: str, rev: str) -> str: + tmp = tempfile.mkdtemp(prefix="json-census-") + proc = subprocess.run( + ["git", "-C", repo, "archive", rev], + stdout=subprocess.PIPE, + check=True, + ) + with tarfile.open(fileobj=io.BytesIO(proc.stdout)) as tar: + tar.extractall(tmp) + return tmp + + +# -------------------------------------------------------------------------- +# The rediscovery gate +# +# Each entry is a defect that was found by hand, at cost, after it reached a +# user. The census must find every one of them again from source alone. The +# entries name a file, the owning type or function, the field and the class -- +# never a line number, so that the gate keeps working when the code moves. +# That matters concretely: two of these surfaced only because a refactor moved +# them into a new file, where a diff-scoped review mistook a pre-existing +# defect for an added line. +# -------------------------------------------------------------------------- + +# Each entry is (name, component, owner, field, class, note). +# +# `component` is a path fragment the matching row's file must contain. It is a +# component, not a file: keying to a file would let exactly the refactor that +# exposed one -- moving code from client.hpp to client.cpp -- turn a +# rediscovery into a silent pass. But it cannot be dropped either. Both the +# client and the server have a `dispatch_response`, and with no component +# constraint the server's reverse-RPC gate matched the client's row and +# reported PASS while covering nothing. +# +# `owner` is the deserialised type, or the enclosing function, or a tuple when +# a defect has more than one manifestation across revisions. +KNOWN_DEFECTS = [ + ( + "progress-total", + "protocol", + "ProgressNotificationParams", + "total", + "null-fragile", + "a present null total ended the client session", + ), + ( + "progress-message", + "protocol", + "ProgressNotificationParams", + "message", + "null-fragile", + "a present null message ended the client session", + ), + ( + "error-message", + "protocol", + "Error", + "message", + "required", + "Error::from_json requires message via at(), so a peer error without one throws", + ), + ( + # Either manifestation counts: at the merge base the reverse-RPC + # correlation path is dispatch_response alone, and the envelope helper + # that carries the contains(result) == contains(error) predicate is + # newer. + "server-reverse-rpc", + "server", + ("is_valid_response_envelope", "dispatch_response"), + "error", + "misroute", + "a response carrying a real result plus 'error': null is rejected and " + "dropped silently on the correlation path, hanging a server-initiated " + "sampling, roots or elicitation request to its timeout", + ), + ( + "client-error", + "client", + "dispatch_response", + "error", + "null-fragile", + "pre-existing at merge base; surfaced only when a refactor moved it", + ), + ( + "client-params", + "client", + "dispatch_notification", + "params", + "null-fragile", + "pre-existing at merge base; surfaced only when a refactor moved it", + ), +] + + +def run_self_test(rows: list[Row], rev: str, stream, expect_fixed: bool = False) -> int: + """Require the census to rediscover every defect that was found by hand. + + Against the merge base every entry must be found: that is the acceptance + criterion, and a miss means the enumeration is incomplete. Against a tree + where the fixes have landed the same run inverts -- a row that is still + found is a fix that did not take. ``expect_fixed`` says which reading + applies, so the exit code means something in both directions instead of + reporting success as failure. + """ + failures = 0 + stream.write(f"rediscovery gate against {rev}\n") + for name, component, owner, field_name, expected, note in KNOWN_DEFECTS: + owners = (owner,) if isinstance(owner, str) else owner + hits = [ + r + for r in rows + if r.field_name == field_name + and r.defect_class == expected + and component in r.file.replace(os.sep, "/") + and ( + owner is None + or any(o in (r.owner, r.function, f"({r.function})") for o in owners) + ) + ] + if hits: + h = hits[0] + stream.write( + f" {'STILL PRESENT' if expect_fixed else 'PASS '} {name:22s} " + f"{h.file}:{h.line} {h.defect_class} fatality={h.fatality}\n" + ) + if expect_fixed: + failures += 1 + elif expect_fixed: + stream.write(f" FIXED {name:22s} no {expected} row remains\n") + else: + failures += 1 + stream.write( + f" FAIL {name:22s} no {expected} row for " + f"{owner or '*'}.{field_name} anywhere under '{component}'\n" + f" ({note})\n" + ) + total = len(KNOWN_DEFECTS) + if expect_fixed: + stream.write(f"fixed {total - failures}/{total}; {failures} still present\n") + else: + stream.write(f"rediscovered {total - failures}/{total}\n") + return failures + check_asymmetry(rows, stream) + + +def check_asymmetry(rows: list[Row], stream) -> int: + """Hold the line between 'error' and 'result'. + + Null-tolerance applies to `"error"` and not to `"result"`: `"result": null` + is a legitimate empty result in JSON-RPC, while `"error": null` is not an + error. A rule that flattens the two closes one defect and opens another, + so the distinction is checked rather than trusted -- including against a + future edit to ENVELOPE_KEYS itself. + """ + failures = 0 + if ENVELOPE_KEYS.get("result", (None,))[0] != "ok": + failures += 1 + stream.write( + " ASYMMETRY FAIL 'result' is not marked ok: a null result is a " + "legitimate empty result and must not be treated as absent\n" + ) + if ENVELOPE_KEYS.get("error", (None,))[0] != MISROUTE: + failures += 1 + stream.write( + " ASYMMETRY FAIL 'error' is not marked misroute: a null error is " + "not an error and a presence test misroutes it\n" + ) + + bad = [r for r in rows if r.field_name == "result" and r.null == MISROUTE] + if bad: + failures += 1 + stream.write(f" ASYMMETRY FAIL {len(bad)} 'result' rows classified misroute:\n") + for r in bad[:10]: + stream.write(f" {r.file}:{r.line} {r.function}\n") + + errors = len([r for r in rows if r.field_name == "error" and r.null == MISROUTE]) + results = len([r for r in rows if r.field_name == "result"]) + if not failures: + stream.write( + f"asymmetry held: {errors} 'error' presence tests flagged misroute, " + f"0 of {results} 'result' rows flagged\n" + ) + return failures + + +# -------------------------------------------------------------------------- +# What this census cannot see +# +# An enumeration whose blind spots are undocumented becomes the next false +# comfort. These are printed by --blind-spots so they travel with the tool. +# -------------------------------------------------------------------------- + +BLIND_SPOTS = [ + ( + "Indirect dispatch is attributed, not traced", + "A handler stored in a std::function and invoked out of a map cannot be " + "followed by name. The census attributes to every peer-receiving lambda " + "the blast radius of the function that dispatches out of a handler " + "registry in the same file. That is right for the registries here, and " + "it would be wrong for a handler invoked from somewhere else.", + ), + ( + "Fatality is static, and public entry points have no caller", + "Server::dispatch is public: an embedder may call it from anywhere, so " + "no static analysis can bound what a throw there costs. Rows reached " + "only that way read 'unknown', which means unproven, not safe.", + ), + ( + "Peer-controlled is asserted for nlohmann::json parameters", + "Every function parameter of json type is treated as peer-controlled. " + "In this codebase each one is reached from json::parse of bytes off the " + "wire, but the census asserts that rather than proving it. A json " + "assembled locally and passed in would be over-reported, never missed.", + ), + ( + "Value domains are recognised only through validate_* calls", + "A field constrained by an inline comparison rather than a validate_* " + "helper reads as unconstrained, so its oversized column will claim an " + "arbitrary value is accepted when it is not.", + ), + ( + "Size limits are looked for lexically, in the same function", + "A ceiling applied by the transport before the document reaches the " + "deserialiser is invisible here. Every row says 'unbounded' because no " + "deserialisation site bounds anything -- that is a statement about " + "these sites, not about the system.", + ), + ( + "Only include/ and src/ are scanned", + "Deserialisation in examples/, conformance/ or benchmark/ is out of " + "scope. So is anything reached through a template the census cannot " + "resolve to a concrete type.", + ), + ( + "The matrix cannot synthesise a baseline for shape-dispatching types", + "std::variant serialisers that choose an arm by inspecting the document " + "are excluded by name with a reason. They are the types where a " + "hand-written case is still required.", + ), +] + + +# -------------------------------------------------------------------------- +# Dispositions +# +# A flagged row that nobody has looked at is not a finding, it is a backlog +# entry pretending to be one. Every flagged row must match a rule in +# scripts/json_census_dispositions.json saying what it means, and the check +# fails if any does not -- so a newly intolerant site cannot arrive without +# someone deciding about it. +# -------------------------------------------------------------------------- + +DISPOSITIONS_PATH = os.path.join( + os.path.dirname(os.path.abspath(__file__)), "json_census_dispositions.json" +) + + +def load_dispositions(path: str) -> list[dict]: + with open(path, encoding="utf-8") as fh: + return json.load(fh)["rules"] + + +def disposition_for(row: Row, rules: list[dict]) -> Optional[dict]: + for rule in rules: + if all( + str(getattr(row, k, None)) == str(v) or (k == "defect_class" and row.defect_class == v) + for k, v in rule["match"].items() + ): + return rule + return None + + +def check_dispositions(rows: list[Row], rules: list[dict], stream) -> int: + flagged = [r for r in rows if r.flagged] + tally: dict[str, int] = {} + unmatched: list[Row] = [] + for r in flagged: + rule = disposition_for(r, rules) + if rule is None: + unmatched.append(r) + continue + key = f"{rule['disposition']}: {rule['reason'][:60]}..." + tally[key] = tally.get(key, 0) + 1 + + stream.write(f"flagged rows: {len(flagged)}\n") + for key, count in sorted(tally.items(), key=lambda kv: -kv[1]): + stream.write(f" {count:4d} {key}\n") + if unmatched: + stream.write(f"\n{len(unmatched)} flagged rows have no disposition:\n") + for r in unmatched[:40]: + stream.write( + f" {r.file}:{r.line} {r.owner}.{r.field_name} " + f"[{r.defect_class}] {r.guard[:50]}\n" + ) + stream.write( + "\nAdd a rule to scripts/json_census_dispositions.json saying whether " + "each is a defect to fix, correct behaviour to accept, or a known " + "cost to record.\n" + ) + return len(unmatched) + + +# -------------------------------------------------------------------------- +# Output +# -------------------------------------------------------------------------- + +COLUMNS = [ + "file", + "line", + "owner", + "field_name", + "site_kind", + "absent", + "null", + "wrong_type", + "oversized", + "guard", + "fatality", + "fatality_basis", + "constrained", + "function", +] + + +def emit(rows: list[Row], fmt: str, stream) -> None: + if fmt == "json": + json.dump([asdict(r) for r in rows], stream, indent=2) + stream.write("\n") + return + if fmt == "csv": + writer = csv.DictWriter(stream, fieldnames=COLUMNS + ["snippet"]) + writer.writeheader() + for r in rows: + writer.writerow({k: getattr(r, k) for k in COLUMNS + ["snippet"]}) + return + widths = {c: max(len(c), *(len(str(getattr(r, c))) for r in rows)) for c in COLUMNS} if rows else {} + stream.write(" ".join(c.ljust(widths[c]) for c in COLUMNS) + "\n") + stream.write(" ".join("-" * widths[c] for c in COLUMNS) + "\n") + for r in rows: + stream.write(" ".join(str(getattr(r, c)).ljust(widths[c]) for c in COLUMNS) + "\n") + + +def build_rows(base: str) -> list[Row]: + """The census of ``base``: every row, with every post-pass applied. + + Both the census CLI and the matrix generator go through here. When the + generator built its rows directly from scan_file() it silently skipped the + target-type and fatality refinements, so the matrix asserted verdicts the + census no longer held -- 54 cases failing against a prediction nothing was + still making. + """ + files = collect_files(base) + index = build_index(files) + + rows: list[Row] = [] + structs: dict[str, list[tuple[str, str]]] = {} + enums: dict[str, str] = {} + for path, rel in files: + file_rows, file_structs = scan_file(path, rel, index) + rows.extend(file_rows) + for name, members in file_structs.items(): + structs.setdefault(name, members) + enums.update( + enum_table(strip_comments(Path(path).read_text(encoding="utf-8", errors="replace"))) + ) + + # Types a peer can steer a document into, and what each delegates to. + deserialisable = {r.owner for r in rows if r.owner and not r.owner.startswith("(")} + delegations: dict[str, set[str]] = {} + for r in rows: + if not r.owner or r.owner.startswith("("): + continue + delegations.setdefault(r.owner, set()).update( + t for t in r.target_types if t in deserialisable + ) + # Macro-defined serialisers reach their members' types through the + # members' declarations rather than through a visible get. + for mtype, mname in structs.get(r.owner, []): + delegations[r.owner].update( + t for t in named_types(mtype) if t in deserialisable + ) + + required_types = {r.owner for r in rows if r.absent == THROW and r.owner} + fill_target_hints(rows, structs) + # A hand-written serialiser that raises on a shape it does not accept -- + # RequestId rejects anything that is not a string or an integer -- is not + # "a type with no required field", even though it reads no field at all. + strict_types = strict_serialisers(collect_files(base)) + refine_targets(rows, enums, required_types, deserialisable, strict_types) + propagate_type_fatality(rows, delegations, deserialisable) + rows.sort(key=lambda r: (r.file, r.line, r.field_name)) + return rows + + +def main(argv: list[str]) -> int: + ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + ap.add_argument("--repo", default=os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + ap.add_argument("--rev", help="census a git revision instead of the working tree") + ap.add_argument("--format", choices=("table", "json", "csv"), default="table") + ap.add_argument("--only-flagged", action="store_true", help="rows that reject some well-formed input") + ap.add_argument("--null-fragile", action="store_true", help="rows that accept absent but reject null") + ap.add_argument( + "--require", + action="append", + default=[], + metavar="FILE:LINE", + help="exit non-zero unless this site is present and flagged", + ) + ap.add_argument("--summary", action="store_true") + ap.add_argument( + "--blind-spots", + action="store_true", + help="print what this census cannot see", + ) + ap.add_argument( + "--check-dispositions", + action="store_true", + help="require every flagged row to have a disposition (exit 1 if any does not)", + ) + ap.add_argument("--dispositions", default=DISPOSITIONS_PATH) + ap.add_argument( + "--fix-list", + action="store_true", + help="print only the rows whose disposition is 'fix'", + ) + ap.add_argument( + "--self-test", + action="store_true", + help="require the census to rediscover every hand-found defect (exit 1 if not)", + ) + ap.add_argument( + "--expect-fixed", + action="store_true", + help="invert --self-test: every known defect must be GONE (exit 1 if one remains)", + ) + args = ap.parse_args(argv) + + if args.blind_spots: + for title, detail in BLIND_SPOTS: + sys.stdout.write(f"{title}\n") + for line in re.findall(r".{1,76}(?:\s|$)", detail): + sys.stdout.write(f" {line.strip()}\n") + sys.stdout.write("\n") + return 0 + + base = export_revision(args.repo, args.rev) if args.rev else args.repo + rows = build_rows(base) + files = collect_files(base) + + rows.sort(key=lambda r: (r.file, r.line, r.field_name)) + + selected = rows + if args.null_fragile: + selected = [r for r in rows if r.absent == OK and r.null in (THROW, MISROUTE)] + elif args.only_flagged: + selected = [r for r in rows if r.flagged] + + status = 0 + rules = load_dispositions(args.dispositions) + if args.check_dispositions: + status |= 1 if check_dispositions(rows, rules, sys.stderr) else 0 + if args.fix_list: + selected = [ + r + for r in rows + if r.flagged + and (disposition_for(r, rules) or {}).get("disposition") == "fix" + ] + if args.self_test: + status |= ( + 1 + if run_self_test( + rows, args.rev or "the working tree", sys.stderr, args.expect_fixed + ) + else 0 + ) + + for req in args.require: + fname, _, lineno = req.rpartition(":") + hits = [ + r for r in rows if r.file.endswith(fname) and r.line == int(lineno) and r.flagged + ] + if hits: + sys.stderr.write( + f"REQUIRE ok {req}: {hits[0].owner}.{hits[0].field_name} " + f"{hits[0].defect_class} fatality={hits[0].fatality}\n" + ) + else: + sys.stderr.write(f"REQUIRE FAIL {req}: not flagged by the census\n") + status = 1 + + if args.summary: + total = len(rows) + flagged = len([r for r in rows if r.flagged]) + fragile = len([r for r in rows if r.absent == OK and r.null in (THROW, MISROUTE)]) + sys.stderr.write( + f"sites={total} flagged={flagged} null-fragile={fragile} " + f"files={len(files)} rev={args.rev or 'working tree'}\n" + ) + + emit(selected, args.format, sys.stdout) + return status + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1:])) diff --git a/scripts/json_census_dispositions.json b/scripts/json_census_dispositions.json new file mode 100644 index 0000000..547a6b7 --- /dev/null +++ b/scripts/json_census_dispositions.json @@ -0,0 +1,102 @@ +{ + "_comment": [ + "A disposition for every row the census flags. 'Flagged but unexamined' is", + "not a terminal state: scripts/json_census.py --check-dispositions fails if", + "any flagged row matches no rule here, so a new intolerant site cannot", + "arrive without someone deciding what it means.", + "", + "Rules are tried in order; the first match wins. Match keys are compared", + "against the row's fields, plus the derived 'defect_class'.", + "", + "Dispositions: 'fix' (a defect, routed to an owner), 'accept' (correct", + "behaviour, with the reason it is correct), 'record' (real but not worth", + "changing now, and known)." + ], + "rules": [ + { + "match": { + "defect_class": "misroute", + "field_name": "error" + }, + "disposition": "fix", + "reason": + "'error': null is not an error. A peer that serialises absent optionals as explicit null -- Go without omitempty, a dumped Python dataclass -- sends this as a matter of course, and a presence test routes the message to the error arm or rejects it outright. On a correlation path the request is then dropped silently and hangs to its timeout.", + "owner": "server/client envelope" + }, + { + "match": { + "defect_class": "misroute", + "field_name": "id" + }, + "disposition": "fix", + "reason": + "'id': null is what JSON-RPC prescribes when the id cannot be determined. Testing presence alone treats that as an id and correlates against a null.", + "owner": "server/client envelope" + }, + { + "match": { + "defect_class": "null-fragile", + "fatality": "session" + }, + "disposition": "fix", + "reason": + "A present-and-null optional throws on a path where the throw ends the session: the read loop leaves its for(;;) and no further message is read. This is the class that has now outrun prediction three times.", + "owner": "client and server dispatch" + }, + { + "match": { + "defect_class": "null-fragile", + "fatality": "process" + }, + "disposition": "fix", + "reason": "A present-and-null optional throws where nothing catches it.", + "owner": "client and server dispatch" + }, + { + "match": { + "defect_class": "null-fragile" + }, + "disposition": "fix", + "reason": + "A peer that serialises absent optionals as explicit null is careless, not hostile, and reaches every one of these. The throw is contained to one message here rather than the session, which makes it lower priority than the session-fatal rows -- not correct behaviour.", + "owner": "protocol serialiser" + }, + { + "match": { + "defect_class": "required", + "site_kind": "macro_required" + }, + "disposition": "accept", + "reason": + "NLOHMANN_DEFINE_TYPE reads every listed member with at(). Those members are non-optional in the struct, so a message missing one is malformed and rejecting it is the correct reading of the protocol.", + "owner": null + }, + { + "match": { + "defect_class": "required" + }, + "disposition": "accept", + "reason": + "The field is read with at() and is declared non-optional: it is required by the protocol, and a message without it is malformed. Rejecting it is correct, provided the throw is contained -- which the fatality column records separately.", + "owner": null + }, + { + "match": { + "defect_class": "untyped" + }, + "disposition": "accept", + "reason": + "A value of the wrong JSON type is malformed under any reading of the protocol. Throwing is correct; only the blast radius is in question, and the fatality column carries that.", + "owner": null + }, + { + "match": { + "oversized": "unbounded" + }, + "disposition": "record", + "reason": + "No deserialisation site bounds a value's size. A ceiling belongs at the transport, where the message is read, rather than at each of several hundred field accesses -- bounding here would be both incomplete and in the wrong place. Recorded so the absence is deliberate and visible.", + "owner": null + } + ] +} diff --git a/scripts/test_check_json_matrix.py b/scripts/test_check_json_matrix.py new file mode 100644 index 0000000..d9fc022 --- /dev/null +++ b/scripts/test_check_json_matrix.py @@ -0,0 +1,288 @@ +#!/usr/bin/env python3 +"""Mutation tests for the peer-input matrix staleness check and its ratchet. + +The real repository is rendered once; every case then mutates the committed +text or manifest in memory and asserts which way the check reads the change. +""" + +from __future__ import annotations + +import json +import os +import sys +import unittest + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + +import check_json_matrix as check # noqa: E402 +import gen_json_matrix as gen # noqa: E402 + +REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +FLAGS = ("absent", "null", "wrong_type", "oversized") + + +class JsonMatrixCheckTest(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + core = os.path.join(REPO, "test", "core") + with open(os.path.join(core, "json_peer_input_matrix_test.cpp"), encoding="utf-8") as fh: + cls.committed = fh.read() + with open(os.path.join(core, "json_matrix_manifest.json"), encoding="utf-8") as fh: + cls.committed_manifest = json.load(fh) + cls.rendered, cls.manifest = gen.render(REPO) + cls.parsed = gen.parse_matrix(cls.rendered) + + def run_check(self, text: str, manifest: dict | None = None) -> tuple[int, str]: + code, lines = check.check( + text, self.manifest if manifest is None else manifest, self.rendered, self.manifest + ) + return code, "\n".join(lines) + + def find_row(self, predicate) -> tuple[str, str]: + for test, block in self.parsed["tests"].items(): + for key, row in block["rows"].items(): + if predicate(row): + return test, key + self.fail("no rendered row matches") + + def block_span(self, text: str, test: str) -> tuple[int, int]: + heads = list(gen._TEST_RE.finditer(text)) + for i, head in enumerate(heads): + if head.group(1) == test: + end = heads[i + 1].start() if i + 1 < len(heads) else len(text) + return head.start(), end + self.fail(f"no TEST {test}") + + def edit_row(self, text: str, test: str, key: str, new_row) -> str: + """Replaces (or with None, removes) one row of one TEST in `text`.""" + start, end = self.block_span(text, test) + block = text[start:end] + for m in gen._ROW_RE.finditer(block): + if m.group(1) != key: + continue + if new_row is None: + # The row and its trailing comma. + block = block[: m.start()] + block[m.end() + 1 :] + else: + flags = ", ".join(str(new_row[f]).lower() for f in FLAGS) + row = f'{{"{key}", R"json({m.group(2)})json", {m.group(3)}, {flags}}}' + block = block[: m.start()] + row + block[m.end() :] + return text[:start] + block + text[end:] + self.fail(f"no row {test}.{key}") + + def flipped(self, test: str, key: str, **cols: bool) -> str: + row = dict(self.parsed["tests"][test]["rows"][key], **cols) + return self.edit_row(self.rendered, test, key, row) + + # 1 + def test_committed_matrix_matches_render(self) -> None: + code, out = self.run_check(self.committed, self.committed_manifest) + self.assertEqual(code, 0, out) + self.assertIn("matches scripts/gen_json_matrix.py", out) + + # 2 + def test_line_wrapping_and_crlf_do_not_count(self) -> None: + wrapped = self.rendered.replace('json::parse(R"json', 'json::parse(\n R"json') + wrapped = wrapped.replace("\n", "\r\n") + self.assertNotEqual(wrapped, self.rendered) + code, out = self.run_check(wrapped) + self.assertEqual(code, 0, out) + self.assertEqual(gen.parse_matrix(wrapped), self.parsed) + + # 3 + def test_toward_tolerance_is_stale(self) -> None: + test, key = self.find_row(lambda r: not r["absent"] and not r["null"]) + code, out = self.run_check(self.flipped(test, key, null=True)) + self.assertEqual(code, 1, out) + self.assertIn(f"TEST(JsonPeerInputMatrix, {test}): {key} null true->false", out) + self.assertIn("decoder fix landed", out) + self.assertNotIn("REGRESSION", out) + + # 4 + def test_optional_member_throwing_on_null_is_a_regression(self) -> None: + test, key = self.find_row(lambda r: not r["absent"] and r["null"]) + committed = self.flipped(test, key, null=False) + # A committed file written before the regression also counted one + # null-fragile row fewer in its header; that alone is not harness drift. + rows = [r for b in self.parsed["tests"].values() for r in b["rows"].values()] + fragile = sum(not r["absent"] and r["null"] for r in rows) + before = f"// {fragile} of the " + self.assertIn(before, committed) + committed = committed.replace(before, f"// {fragile - 1} of the ", 1) + code, out = self.run_check(committed) + self.assertEqual(code, 2, out) + self.assertIn("REGRESSION", out) + self.assertIn(f"{test}): {key} null false->true (optional member", out) + self.assertIn("treat an explicit null as absent", out) + self.assertNotIn("became required", out) + self.assertNotIn("harness text", out) + + # 5 + def test_member_becoming_required_is_a_regression(self) -> None: + test, key = self.find_row(lambda r: r["absent"]) + code, out = self.run_check(self.flipped(test, key, absent=False)) + self.assertEqual(code, 2, out) + self.assertIn(f"{test}): {key} absent false->true (member is now required)", out) + self.assertIn("became required", out) + self.assertNotIn("treat an explicit null", out) + + # 6 + def test_stricter_type_or_domain_validation_is_stale(self) -> None: + for col in ("wrong_type", "oversized"): + test, key = self.find_row(lambda r, c=col: r[c]) + code, out = self.run_check(self.flipped(test, key, **{col: False})) + self.assertEqual(code, 1, out) + self.assertIn(f"{key} {col} false->true (stricter validation)", out) + self.assertNotIn("REGRESSION", out) + + # 7 + def test_new_rows_are_judged_by_their_null_column(self) -> None: + test, key = self.find_row(lambda r: not r["absent"] and r["null"]) + code, out = self.run_check(self.edit_row(self.rendered, test, key, None)) + self.assertEqual(code, 2, out) + self.assertIn(f"{test}): field '{key}' added as null-fragile", out) + + test, key = self.find_row(lambda r: not r["null"]) + code, out = self.run_check(self.edit_row(self.rendered, test, key, None)) + self.assertEqual(code, 1, out) + self.assertIn(f"{test}): field '{key}' added", out) + + # 8 + def test_missing_test_is_stale(self) -> None: + # A TEST with no null-fragile row, so dropping it is only stale. + test = next( + t + for t, block in self.parsed["tests"].items() + if block["rows"] and all(r["absent"] or not r["null"] for r in block["rows"].values()) + ) + start, end = self.block_span(self.rendered, test) + code, out = self.run_check(self.rendered[:start] + self.rendered[end:]) + self.assertEqual(code, 1, out) + self.assertIn(f"TEST(JsonPeerInputMatrix, {test}): only in the generated file", out) + + # 9 + def test_harness_edit_is_stale(self) -> None: + edited = self.rendered.replace('SCOPED_TRACE("null");', 'SCOPED_TRACE("nul");', 1) + self.assertNotEqual(edited, self.rendered) + code, out = self.run_check(edited) + self.assertEqual(code, 1, out) + self.assertIn("harness text", out) + self.assertIn("Review the diff", out) + self.assertNotIn("decoder fix landed", out) + + # Still named when a row difference is reported alongside it. + test, key = self.find_row(lambda r: not r["absent"] and not r["null"]) + row = dict(self.parsed["tests"][test]["rows"][key], null=True) + both = self.edit_row(edited, test, key, row) + code, out = self.run_check(both) + self.assertEqual(code, 1, out) + self.assertIn(f"{key} null true->false", out) + self.assertIn("harness text", out) + + # 10 + def test_manifest_difference_is_stale(self) -> None: + manifest = dict(self.manifest, case_count=self.manifest["case_count"] + 4) + code, out = self.run_check(self.rendered, manifest) + self.assertEqual(code, 1, out) + self.assertIn("manifest key 'case_count'", out) + + # 11 + def test_generator_refuses_unaccepted_regressions(self) -> None: + test, key = self.find_row(lambda r: not r["absent"] and r["null"]) + committed = self.flipped(test, key, null=False) + + code, lines = gen.guard([committed], self.rendered, []) + self.assertEqual(code, 2) + self.assertTrue(any(f"{test}): {key} null false->true" in line for line in lines)) + + for accepted in (f"{test}.{key}", test): + code, lines = gen.guard([committed], self.rendered, [accepted]) + self.assertEqual((code, lines), (0, []), accepted) + + code, lines = gen.guard([self.rendered], self.rendered, ["NoSuchType.key"]) + self.assertEqual(code, 1) + self.assertEqual(lines, ["--accept-regression NoSuchType.key matches no regressed row"]) + + # A nested type may be named as written in C++, with :: for the _. + test, key = next( + (t, k) + for t, block in self.parsed["tests"].items() + if "_" in t + for k, r in block["rows"].items() + if not r["absent"] and r["null"] + ) + committed = self.flipped(test, key, null=False) + code, lines = gen.guard([committed], self.rendered, [test.replace("_", "::") + "." + key]) + self.assertEqual((code, lines), (0, [])) + + # 13 + def test_required_member_made_optional_but_null_fragile_is_a_regression(self) -> None: + test, key = self.find_row(lambda r: not r["absent"] and r["null"]) + committed = self.flipped(test, key, absent=True) + code, out = self.run_check(committed) + self.assertEqual(code, 2, out) + self.assertIn(f"{test}): {key} became optional but throws on explicit null", out) + self.assertNotIn("now tolerated", out) + self.assertEqual(gen.guard([committed], self.rendered, [])[0], 2) + + # 14 + def test_guard_checks_head_as_well_as_the_working_tree(self) -> None: + test, key = self.find_row(lambda r: not r["absent"] and r["null"]) + head = self.flipped(test, key, null=False) + # Emptying or deleting the output file must not leave nothing to compare. + for out_text in ("", None, "// truncated\n"): + baselines = gen.guard_baselines(out_text, head) + self.assertEqual(baselines, [head]) + self.assertEqual(gen.guard(baselines, self.rendered, [])[0], 2) + self.assertEqual(gen.guard_baselines(self.rendered, head), [head, self.rendered]) + self.assertEqual(gen.guard_baselines(head, head), [head]) + self.assertEqual(gen.guard_baselines(self.rendered, None), [self.rendered]) + self.assertEqual(gen.guard_baselines("", None), []) + self.assertEqual(gen.guard_baselines(None, ""), []) + + # 15 + def test_row_deleted_from_the_working_tree_but_at_head_is_still_a_regression(self) -> None: + # A required row with a tolerant past, in a TEST with no null-fragile + # row, so a working copy without the row or the TEST reads as stale. + test, key = next( + (t, k) + for t, block in self.parsed["tests"].items() + if all(r["absent"] or not r["null"] for r in block["rows"].values()) + for k, r in block["rows"].items() + if r["absent"] and not r["null"] + ) + head = self.flipped(test, key, absent=False) + start, end = self.block_span(head, test) + for working in (self.edit_row(head, test, key, None), head[:start] + head[end:]): + self.assertEqual(gen.guard([working], self.rendered, [])[0], 0) + code, lines = gen.guard(gen.guard_baselines(working, head), self.rendered, []) + self.assertEqual(code, 2, lines) + self.assertTrue(any(f"{test}): {key} absent false->true" in line for line in lines)) + baselines = gen.guard_baselines(working, head) + self.assertEqual(gen.guard(baselines, self.rendered, [f"{test}.{key}"]), (0, [])) + + # 12 + def test_every_emitted_line_is_ascii(self) -> None: + fragile = self.find_row(lambda r: not r["absent"] and r["null"]) + tolerant = self.find_row(lambda r: not r["absent"] and not r["null"]) + regressed = self.flipped(*fragile, null=False) + tolerant_row = dict(self.parsed["tests"][tolerant[0]]["rows"][tolerant[1]], null=True) + both = self.edit_row(regressed, *tolerant, tolerant_row) + harness = self.rendered.replace('SCOPED_TRACE("absent");', 'SCOPED_TRACE("gone");', 1) + manifest = dict(self.manifest, mode_count=5) + emitted: list[str] = [] + for text, man in ( + (self.committed, self.committed_manifest), + (both, manifest), + (harness, self.manifest), + ("no tests here", self.manifest), + ): + emitted += check.check(text, man, self.rendered, self.manifest)[1] + emitted += gen.guard([regressed], self.rendered, ["NoSuchType"])[1] + self.assertIn("REGRESSION", "\n".join(emitted)) + self.assertIn("is stale", "\n".join(emitted)) + for line in emitted: + self.assertTrue(line.isascii(), line) + +if __name__ == "__main__": + unittest.main() diff --git a/scripts/test_release_workflow_policy.py b/scripts/test_release_workflow_policy.py index f85e224..1ab010e 100644 --- a/scripts/test_release_workflow_policy.py +++ b/scripts/test_release_workflow_policy.py @@ -82,6 +82,14 @@ def test_rejects_publisher_without_anchor_handoff(self) -> None: blocks["aur"] = [line.replace("Verify fixed anchor handoff", "Trust downloaded files") for line in blocks["aur"]] self.assert_policy_error(lambda: policy.check_publication_invariants(blocks, self.text)) + def test_rejects_aur_without_passphrase_protection(self) -> None: + blocks = self.mutated_blocks("SSH_ASKPASS_REQUIRE=force", "SSH_ASKPASS_REQUIRE=never") + self.assert_policy_error(lambda: policy.check_publication_invariants(blocks, self.text)) + + def test_rejects_aur_without_batch_mode(self) -> None: + blocks = self.mutated_blocks("BatchMode=yes", "BatchMode=no") + self.assert_policy_error(lambda: policy.check_publication_invariants(blocks, self.text)) + def test_rejects_unpinned_cloudsmith_cli(self) -> None: mutated = self.text.replace('cli-version: "1.19.0"', 'cli-version: "latest"', 1) blocks = policy.job_blocks(mutated.splitlines()) diff --git a/src/auth/challenge.cpp b/src/auth/challenge.cpp new file mode 100644 index 0000000..1f4b6ae --- /dev/null +++ b/src/auth/challenge.cpp @@ -0,0 +1,316 @@ +#include + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mcp::auth { + +namespace { + +constexpr std::string_view g_token_specials = "!#$%&'*+-.^_`|~"; +constexpr int g_hex_base = 16; + +bool is_token_char(unsigned char character) { + return std::isalnum(character) != 0 || + g_token_specials.find(static_cast(character)) != std::string_view::npos; +} + +std::string to_lower_ascii(std::string value) { + std::transform(value.begin(), value.end(), value.begin(), + [](unsigned char character) { return static_cast(std::tolower(character)); }); + return value; +} + +bool equals_ignore_case(std::string_view left, std::string_view right) { + return left.size() == right.size() && + std::equal(left.begin(), left.end(), right.begin(), [](char lhs, char rhs) { + return std::tolower(static_cast(lhs)) == + std::tolower(static_cast(rhs)); + }); +} + +void skip_whitespace(std::string_view text, std::size_t& position) { + while (position < text.size() && (text[position] == ' ' || text[position] == '\t')) { + ++position; + } +} + +void skip_separators(std::string_view text, std::size_t& position) { + while (position < text.size() && + (text[position] == ' ' || text[position] == '\t' || text[position] == ',')) { + ++position; + } +} + +std::string read_token(std::string_view text, std::size_t& position) { + const auto start = position; + while (position < text.size() && is_token_char(static_cast(text[position]))) { + ++position; + } + return std::string(text.substr(start, position - start)); +} + +/// Read a quoted-string, expanding backslash escapes. Returns nullopt when unterminated. +std::optional read_quoted_string(std::string_view text, std::size_t& position) { + ++position; + std::string value; + while (position < text.size()) { + const char character = text[position++]; + if (character == '\\') { + if (position >= text.size()) { + return std::nullopt; + } + value.push_back(text[position++]); + } else if (character == '"') { + return value; + } else { + value.push_back(character); + } + } + return std::nullopt; +} + +/// Record a parameter, keeping the first occurrence of any repeated name. +void assign_parameter(BearerChallenge& challenge, std::string name, std::string value) { + const auto assign_once = [&value](std::optional& field) { + if (!field) { + field = value; + } + }; + if (name == "realm") { + assign_once(challenge.realm); + } else if (name == "resource_metadata") { + assign_once(challenge.resource_metadata); + } else if (name == "scope") { + assign_once(challenge.scope); + } else if (name == "error") { + assign_once(challenge.error); + } else if (name == "error_description") { + assign_once(challenge.error_description); + } else if (name == "error_uri") { + assign_once(challenge.error_uri); + } + challenge.parameters.emplace_back(std::move(name), std::move(value)); +} + +/// Decode one `application/x-www-form-urlencoded` component. +std::string form_decode(std::string_view value) { + std::string result; + result.reserve(value.size()); + for (std::size_t index = 0; index < value.size(); ++index) { + const char character = value[index]; + if (character == '+') { + result.push_back(' '); + } else if (character == '%' && index + 2 < value.size()) { + const auto high = mcp::constants::g_hex_digits_upper.find( + static_cast(std::toupper(static_cast(value[index + 1])))); + const auto low = mcp::constants::g_hex_digits_upper.find( + static_cast(std::toupper(static_cast(value[index + 2])))); + if (high == std::string_view::npos || low == std::string_view::npos) { + result.push_back(character); + continue; + } + result.push_back(static_cast((high * g_hex_base) + low)); + index += 2; + } else { + result.push_back(character); + } + } + return result; +} + +void assign_response_field(AuthorizationResponse& response, const std::string& name, + std::string value) { + const auto assign_once = [&value](std::optional& field) { + if (!field) { + field = value; + } + }; + if (name == "code") { + assign_once(response.code); + } else if (name == "state") { + assign_once(response.state); + } else if (name == "iss") { + assign_once(response.iss); + } else if (name == "error") { + assign_once(response.error); + } else if (name == "error_description") { + assign_once(response.error_description); + } else if (name == "error_uri") { + assign_once(response.error_uri); + } +} + +} // namespace + +bool BearerChallenge::is_bearer() const { return equals_ignore_case(scheme, "Bearer"); } + +std::vector parse_www_authenticate(std::string_view header_value) { + std::vector challenges; + std::optional current; + + const auto flush = [&challenges, ¤t]() { + if (current) { + challenges.push_back(std::move(*current)); + current.reset(); + } + }; + + std::size_t position = 0; + skip_separators(header_value, position); + while (position < header_value.size()) { + auto name = read_token(header_value, position); + if (name.empty()) { + break; // Not a token where one is required; discard the remainder. + } + + const auto after_name = position; + skip_whitespace(header_value, position); + if (position >= header_value.size() || header_value[position] != '=') { + // A bare token introduces the next challenge rather than a parameter. + position = after_name; + flush(); + current.emplace(); + current->scheme = std::move(name); + skip_separators(header_value, position); + continue; + } + + ++position; + // `token68` credentials may carry `=` padding; that belongs to the preceding scheme. + auto padding = position; + while (padding < header_value.size() && header_value[padding] == '=') { + ++padding; + } + auto probe = padding; + skip_whitespace(header_value, probe); + if (probe >= header_value.size() || header_value[probe] == ',') { + position = probe; + skip_separators(header_value, position); + continue; + } + + skip_whitespace(header_value, position); + std::string value; + if (header_value[position] == '"') { + auto quoted = read_quoted_string(header_value, position); + if (!quoted) { + break; // Unterminated quoted-string; nothing after it can be trusted. + } + value = std::move(*quoted); + } else { + value = read_token(header_value, position); + } + + if (current) { + assign_parameter(*current, to_lower_ascii(std::move(name)), std::move(value)); + } + skip_separators(header_value, position); + } + + flush(); + return challenges; +} + +std::vector parse_www_authenticate(const std::vector& header_values) { + std::vector challenges; + for (const auto& header_value : header_values) { + auto parsed = parse_www_authenticate(header_value); + challenges.insert(challenges.end(), std::make_move_iterator(parsed.begin()), + std::make_move_iterator(parsed.end())); + } + return challenges; +} + +std::optional select_bearer_challenge(const std::vector& challenges) { + const auto match = + std::find_if(challenges.begin(), challenges.end(), + [](const BearerChallenge& challenge) { return challenge.is_bearer(); }); + if (match == challenges.end()) { + return std::nullopt; + } + return *match; +} + +AuthorizationResponse parse_authorization_response(const std::string& redirect_url) { + AuthorizationResponse response; + + auto query = std::string_view(redirect_url); + if (const auto question = query.find('?'); question != std::string_view::npos) { + query.remove_prefix(question + 1); + } + if (const auto fragment = query.find('#'); fragment != std::string_view::npos) { + query = query.substr(0, fragment); + } + + while (!query.empty()) { + auto field = query; + if (const auto separator = query.find('&'); separator != std::string_view::npos) { + field = query.substr(0, separator); + query.remove_prefix(separator + 1); + } else { + query = {}; + } + + const auto equals = field.find('='); + if (equals == std::string_view::npos) { + continue; + } + assign_response_field(response, to_lower_ascii(form_decode(field.substr(0, equals))), + form_decode(field.substr(equals + 1))); + } + return response; +} + +AuthorizationResponseValidation validate_authorization_response(const AuthorizationRequest& request, + const AuthorizationResponse& response) { + if (!request.state.empty()) { + if (!response.state) { + return {AuthorizationResponseStatus::state_missing, + "authorization response omitted the state parameter"}; + } + if (*response.state != request.state) { + return {AuthorizationResponseStatus::state_mismatch, + "authorization response state did not match the recorded value"}; + } + } + + // RFC 9207 Section 2.4 as adopted by MCP. Plain string comparison only: no case folding, no + // default-port elision, no trailing-slash or percent-encoding normalization. + if (response.iss) { + if (*response.iss != request.issuer) { + return {AuthorizationResponseStatus::issuer_mismatch, + "authorization response iss did not match the recorded issuer"}; + } + } else if (request.issuer_parameter_supported) { + return {AuthorizationResponseStatus::issuer_missing, + "authorization server advertises iss support but omitted the parameter"}; + } + + // Only an issuer-authentic response may have its error values acted on or displayed. + if (response.error) { + // Sanitized at the source rather than where the message is interpolated into an exception, + // so every consumer of this validation message gets the flattened form. The value reached + // us through a redirect the authorization server controls, and being issuer-authentic makes + // it trustworthy as to origin, not as to content. + return { + AuthorizationResponseStatus::server_error, + "authorization server returned error " + detail::sanitize_for_diagnostics(*response.error)}; + } + if (!response.code || response.code->empty()) { + return {AuthorizationResponseStatus::code_missing, + "authorization response carried no authorization code"}; + } + return {AuthorizationResponseStatus::accepted, {}}; +} + +} // namespace mcp::auth diff --git a/src/auth/client_identity.cpp b/src/auth/client_identity.cpp new file mode 100644 index 0000000..35f0b6e --- /dev/null +++ b/src/auth/client_identity.cpp @@ -0,0 +1,237 @@ +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mcp::auth { + +namespace { + +constexpr std::string_view g_offline_access = "offline_access"; + +/// Split a space-delimited scope string into its individual values. +std::vector split_scope(const std::string& scope) { + std::vector values; + std::istringstream stream(scope); + std::string value; + while (stream >> value) { + values.push_back(value); + } + return values; +} + +std::string join_scope(const std::vector& values) { + std::string joined; + for (const auto& value : values) { + if (!joined.empty()) { + joined.push_back(' '); + } + joined += value; + } + return joined; +} + +bool contains(const std::vector& values, std::string_view needle) { + return std::find(values.begin(), values.end(), needle) != values.end(); +} + +} // namespace + +std::string_view describe(ClientIdentitySource source) { + switch (source) { + case ClientIdentitySource::pre_registered: + return "pre-registered credentials"; + case ClientIdentitySource::client_id_metadata_document: + return "client ID metadata document"; + case ClientIdentitySource::dynamic_registration: + return "dynamic client registration"; + } + return "unknown client identity source"; +} + +std::string_view describe(ClientIdentityDecision decision) { + switch (decision) { + case ClientIdentityDecision::use_pre_registered: + return "present the application's pre-registered credentials"; + case ClientIdentityDecision::use_client_id_metadata_document: + return "present the configured client ID metadata document URL"; + case ClientIdentityDecision::reuse_stored_registration: + return "present the credentials already registered with this issuer"; + case ClientIdentityDecision::register_dynamically: + return "register dynamically with this issuer"; + case ClientIdentityDecision::unavailable: + return "no client identity is available for this issuer"; + } + return "unknown client identity decision"; +} + +void to_json(nlohmann::json& json, const OAuthClientMetadata& metadata) { + json = nlohmann::json::object(); + json["redirect_uris"] = metadata.redirect_uris; + // SEP-837: the registration body always states what kind of application is registering. + json["application_type"] = metadata.application_type; + json["grant_types"] = metadata.grant_types; + json["response_types"] = metadata.response_types; + if (metadata.client_name) { + json["client_name"] = *metadata.client_name; + } + if (metadata.client_uri) { + json["client_uri"] = *metadata.client_uri; + } + if (metadata.software_id) { + json["software_id"] = *metadata.software_id; + } + if (metadata.software_version) { + json["software_version"] = *metadata.software_version; + } + if (metadata.scope) { + json["scope"] = *metadata.scope; + } + if (metadata.token_endpoint_auth_method) { + json["token_endpoint_auth_method"] = *metadata.token_endpoint_auth_method; + } +} + +bool OAuthClientInformation::secret_expired(std::int64_t now_seconds) const { + // RFC 7591: zero means the secret never expires; absent means the server said nothing. + if (!client_secret_expires_at || *client_secret_expires_at == 0) { + return false; + } + return *client_secret_expires_at <= now_seconds; +} + +void from_json(const nlohmann::json& json, OAuthClientInformation& information) { + if (json.contains("client_id")) { + json.at("client_id").get_to(information.client_id); + } + if (json.contains("client_secret") && !json.at("client_secret").is_null()) { + information.client_secret = json.at("client_secret").get(); + } + if (json.contains("client_id_issued_at") && json.at("client_id_issued_at").is_number()) { + information.client_id_issued_at = json.at("client_id_issued_at").get(); + } + if (json.contains("client_secret_expires_at") && json.at("client_secret_expires_at").is_number()) { + information.client_secret_expires_at = json.at("client_secret_expires_at").get(); + } +} + +void to_json(nlohmann::json& json, const OAuthClientInformation& information) { + json = nlohmann::json::object(); + json["client_id"] = information.client_id; + if (information.client_secret) { + json["client_secret"] = *information.client_secret; + } + if (information.client_id_issued_at) { + json["client_id_issued_at"] = *information.client_id_issued_at; + } + if (information.client_secret_expires_at) { + json["client_secret_expires_at"] = *information.client_secret_expires_at; + } + json["issuer"] = information.issuer; +} + +struct InMemoryClientCredentialStore::Impl { + mutable std::mutex mutex; + std::unordered_map credentials; +}; + +InMemoryClientCredentialStore::InMemoryClientCredentialStore() : impl_(std::make_unique()) {} + +InMemoryClientCredentialStore::~InMemoryClientCredentialStore() = default; + +void InMemoryClientCredentialStore::store(const std::string& issuer, + OAuthClientInformation information) { + std::lock_guard lock(impl_->mutex); + impl_->credentials[issuer] = std::move(information); +} + +std::optional InMemoryClientCredentialStore::load( + const std::string& issuer) const { + std::lock_guard lock(impl_->mutex); + const auto iter = impl_->credentials.find(issuer); + if (iter == impl_->credentials.end()) { + return std::nullopt; + } + return iter->second; +} + +void InMemoryClientCredentialStore::remove(const std::string& issuer) { + std::lock_guard lock(impl_->mutex); + impl_->credentials.erase(issuer); +} + +ClientIdentityDecision select_client_identity(const ClientIdentityConfig& config, + const ClientIdentityServerFacts& server, + const std::optional& stored) { + // Injected credentials are terminal. Falling back to registration here would swap the client + // the application chose for one the authorization server minted, without telling anybody. + if (config.pre_registered && !config.pre_registered->client_id.empty()) { + // Credentials that name their issuer are bound to it: presenting the secret the application + // holds for one authorization server to a different one is credential misbinding, and the + // terminal rule above means the answer is "no identity", never a silent registration. + const auto& expected_issuer = config.pre_registered->issuer; + if (expected_issuer.empty()) { + // Unbound credentials. A `client_id` is a public identifier and costs nothing to show + // to the wrong server, so a public client still authorizes normally. A `client_secret` + // is different: the authorization server reached here was named by the + // protected-resource document, so presenting an unbound secret hands the application's + // credential to whichever server that document chose. "Bound to no issuer" must not be + // read as "bound to every issuer", so the secret is refused at the point of use. + if (config.pre_registered->client_secret && + !config.pre_registered->client_secret->empty()) { + return ClientIdentityDecision::unavailable; + } + return ClientIdentityDecision::use_pre_registered; + } + if (expected_issuer != server.issuer) { + return ClientIdentityDecision::unavailable; + } + return ClientIdentityDecision::use_pre_registered; + } + + if (config.client_metadata_url && !config.client_metadata_url->empty() && + server.client_id_metadata_document_supported) { + return ClientIdentityDecision::use_client_id_metadata_document; + } + + // A stored entry is only usable when its recorded issuer is the issuer being contacted. An + // authorization server change therefore falls straight through to a fresh registration. + if (stored && !stored->client_id.empty() && stored->issuer == server.issuer && + !server.issuer.empty()) { + return ClientIdentityDecision::reuse_stored_registration; + } + + if (server.registration_endpoint && !server.registration_endpoint->empty()) { + return ClientIdentityDecision::register_dynamically; + } + + return ClientIdentityDecision::unavailable; +} + +nlohmann::json build_registration_request(const OAuthClientMetadata& metadata, + const ClientIdentityServerFacts& server) { + auto body = nlohmann::json(metadata); + + // Offline access is only requested when the server published it; asking for an unpublished + // scope invites an outright rejection of the whole registration. + if (contains(server.scopes_supported, g_offline_access)) { + auto scopes = metadata.scope ? split_scope(*metadata.scope) : std::vector{}; + if (!contains(scopes, g_offline_access)) { + scopes.emplace_back(g_offline_access); + } + body["scope"] = join_scope(scopes); + } + return body; +} + +} // namespace mcp::auth diff --git a/src/auth/metadata_policy.cpp b/src/auth/metadata_policy.cpp new file mode 100644 index 0000000..8b1bbe1 --- /dev/null +++ b/src/auth/metadata_policy.cpp @@ -0,0 +1,574 @@ +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mcp::auth { + +namespace net = boost::asio; + +namespace { + +constexpr std::uint32_t g_octet_mask = 0xFFU; +constexpr int g_octet1_shift = 24; +constexpr int g_octet2_shift = 16; +constexpr int g_octet3_shift = 8; + +/// Decompose an absolute URL into its scheme and authority without normalizing either. +bool split_url(const std::string& url, std::string& scheme, std::string& authority) { + const auto scheme_end = url.find("://"); + if (scheme_end == std::string::npos || scheme_end == 0) { + return false; + } + scheme = url.substr(0, scheme_end); + + const auto authority_start = scheme_end + 3; + auto authority_end = url.size(); + for (auto index = authority_start; index < url.size(); ++index) { + const char character = url[index]; + if (character == '/' || character == '?' || character == '#') { + authority_end = index; + break; + } + } + authority = url.substr(authority_start, authority_end - authority_start); + return !authority.empty(); +} + +/// Extract the host from an authority. Userinfo is rejected by the caller, not stripped here. +bool authority_host(const std::string& authority, std::string& host) { + if (authority.front() == '[') { + const auto closing = authority.find(']'); + if (closing == std::string::npos || closing == 1) { + return false; + } + host = authority.substr(1, closing - 1); + return true; + } + const auto colon = authority.find(':'); + if (colon == std::string::npos) { + host = authority; + return true; + } + if (authority.find(':', colon + 1) != std::string::npos) { + return false; + } + host = authority.substr(0, colon); + return !host.empty(); +} + +/// Lowercase only ASCII letters; every other byte (digits, punctuation, non-ASCII) is left as-is. +std::string lowercase_ascii(std::string value) { + std::transform(value.begin(), value.end(), value.begin(), [](unsigned char ch) -> char { + return (ch >= 'A' && ch <= 'Z') ? static_cast(ch - 'A' + 'a') : static_cast(ch); + }); + return value; +} + +/// Parse a URL port substring strictly: ASCII digits only, no sign, no whitespace, and in range +/// [0, 65535]. Leading zeros are tolerated and simply absorbed into the numeric value (`"00443"` and +/// `"0443"` both parse to 443), but anything that is not a plain unsigned decimal number — including +/// an empty string, a sign, embedded whitespace, or a value that overflows a 16-bit port — returns +/// `nullopt` rather than silently truncating or wrapping. +std::optional parse_strict_port(const std::string& port_text) { + if (port_text.empty()) { + return std::nullopt; + } + std::uint64_t value = 0; + for (const char character : port_text) { + if (character < '0' || character > '9') { + return std::nullopt; + } + value = (value * 10) + static_cast(character - '0'); + if (value > 65535U) { + return std::nullopt; + } + } + return static_cast(value); +} + +/// Canonical form of an origin's components. Kept split (rather than only the formatted string) so +/// `validate_metadata_url` can feed the same canonical scheme and host into the checks that run after +/// the origin decision, instead of re-deriving them from raw, non-canonical text. +struct CanonicalOrigin { + std::string scheme; + std::string host; ///< Never bracketed, even when the origin names an IPv6 literal. + std::optional + port; ///< Absent when no port was written, or it was the scheme's default. + bool bracketed; ///< Whether the original authority wrote the host as `[...]`. +}; + +/// Decompose and canonicalize an origin (`scheme://authority`) for comparison. +/// +/// Lowercases the scheme and a non-IP-literal host, and strips every trailing dot from such a host. A +/// host that parses as an IP literal (IPv4, or IPv6 with or without brackets) is replaced by its +/// normalized textual form, so expanded and compressed IPv6 spellings compare equal; a literal is +/// never dot-stripped. An explicit port is parsed strictly (see `parse_strict_port`) and dropped when +/// it equals the scheme's default (`443` for `https`, `80` for `http`). +/// +/// Returns `nullopt` when `origin` is not `scheme://authority` with nothing following (a path, query +/// or fragment), when the host is malformed or empty after dot-stripping, or when a port is not a +/// plain in-range decimal number. +std::optional canonicalize_origin_parts(const std::string& origin) { + std::string scheme; + std::string authority; + if (!split_url(origin, scheme, authority)) { + return std::nullopt; + } + // A bare origin has nothing past the authority. An entry carrying a path, query or fragment + // (e.g. an allow-list entry mistakenly written as "https://as.test/realms/foo") is not an origin + // and must never be silently truncated into a match for the origin it happens to prefix. + if (scheme.size() + 3 + authority.size() != origin.size()) { + return std::nullopt; + } + std::string host; + if (!authority_host(authority, host)) { + return std::nullopt; + } + + const bool bracketed = authority.front() == '['; + std::string port_text; + bool has_port = false; + if (bracketed) { + const auto closing = authority.find(']'); + if (closing != std::string::npos && closing + 1 < authority.size() && + authority[closing + 1] == ':') { + port_text = authority.substr(closing + 2); + has_port = true; + } + } else { + const auto colon = authority.find(':'); + if (colon != std::string::npos) { + port_text = authority.substr(colon + 1); + has_port = true; + } + } + + std::optional port; + if (has_port) { + port = parse_strict_port(port_text); + if (!port) { + return std::nullopt; + } + } + + auto canonical_scheme = lowercase_ascii(scheme); + + boost::system::error_code error; + const auto address = net::ip::make_address(host, error); + const bool is_ip_literal = !error; + + std::string canonical_host; + if (is_ip_literal) { + canonical_host = address.to_string(); + } else { + canonical_host = lowercase_ascii(host); + while (!canonical_host.empty() && canonical_host.back() == '.') { + canonical_host.pop_back(); + } + if (canonical_host.empty()) { + return std::nullopt; + } + } + + if (port && ((*port == 443U && canonical_scheme == "https") || + (*port == 80U && canonical_scheme == "http"))) { + port.reset(); + } + + return CanonicalOrigin{std::move(canonical_scheme), std::move(canonical_host), port, bracketed}; +} + +/// Format a `CanonicalOrigin` back into `scheme://authority` text, for list comparison and for the +/// value handed to `origin_allowance`. +std::string format_canonical_origin(const CanonicalOrigin& parts) { + std::string authority = parts.bracketed ? "[" + parts.host + "]" : parts.host; + if (parts.port) { + authority += ":" + std::to_string(*parts.port); + } + return parts.scheme + "://" + authority; +} + +/// Canonicalize an origin purely for list-membership comparison; see `canonicalize_origin_parts`. +std::optional canonicalize_origin(const std::string& origin) { + const auto parts = canonicalize_origin_parts(origin); + if (!parts) { + return std::nullopt; + } + return format_canonical_origin(*parts); +} + +/// Compare a canonicalized origin against the allow list, canonicalizing each entry at comparison +/// time so the caller never has to keep a normalized copy of the policy around. An entry that does +/// not canonicalize to a bare origin (a malformed port, a path/query/fragment, or any other +/// malformed form) matches nothing, rather than being widened or truncated into a match for the +/// origin it happens to prefix. Dropping an allow entry grants nothing, so this direction is +/// fail-closed; the deny list, where the same silence would be fail-open, is handled by +/// `denies_origin` instead. +bool contains_origin(const std::vector& origins, const std::string& canonical_origin) { + return std::any_of(origins.begin(), origins.end(), [&](const std::string& candidate) { + const auto canonical_candidate = canonicalize_origin(candidate); + return canonical_candidate && *canonical_candidate == canonical_origin; + }); +} + +/// Compare a canonicalized origin against the deny list, requiring every entry to be a bare origin. +/// +/// An entry that does not canonicalize is a configuration error and is reported, not discarded (which +/// would be fail-open) or reinterpreted as the origin it resembles. The whole list is examined before +/// a match is returned, so a malformed entry is reported even when an earlier entry matched; a policy +/// carrying one refuses every target until it is corrected. +/// +/// @throws MetadataPolicyError With `denied_origin_entry_malformed`, naming the offending entry. +bool denies_origin(const std::vector& origins, const std::string& canonical_origin) { + bool denied = false; + for (const auto& candidate : origins) { + const auto canonical_candidate = canonicalize_origin(candidate); + if (!canonical_candidate) { + throw MetadataPolicyError(MetadataUrlDecision::denied_origin_entry_malformed, candidate); + } + if (*canonical_candidate == canonical_origin) { + denied = true; + } + } + return denied; +} + +/// Classify an IPv4 address against the ranges that must never be reached by a metadata fetch. +MetadataUrlDecision classify_v4(const MetadataFetchPolicy& policy, const net::ip::address_v4& address) { + const auto value = address.to_uint(); + const auto octet1 = static_cast(value >> g_octet1_shift) & g_octet_mask; + const auto octet2 = static_cast(value >> g_octet2_shift) & g_octet_mask; + const auto octet3 = static_cast(value >> g_octet3_shift) & g_octet_mask; + + if (octet1 == 127) { + return policy.allow_plain_http_loopback ? MetadataUrlDecision::allowed + : MetadataUrlDecision::address_loopback; + } + if (octet1 == 169 && octet2 == 254) { + return MetadataUrlDecision::address_link_local; + } + if (octet1 == 10 || (octet1 == 172 && octet2 >= 16 && octet2 <= 31) || + (octet1 == 192 && octet2 == 168)) { + return MetadataUrlDecision::address_private; + } + if (value == 0xFFFFFFFFU) { + return MetadataUrlDecision::address_multicast; + } + if (octet1 >= 224 && octet1 <= 239) { + return MetadataUrlDecision::address_multicast; + } + if (octet1 == 0 || octet1 >= 240) { + return MetadataUrlDecision::address_reserved; + } + if (octet1 == 100 && octet2 >= 64 && octet2 <= 127) { + return MetadataUrlDecision::address_reserved; + } + if (octet1 == 192 && octet2 == 0 && (octet3 == 0 || octet3 == 2)) { + return MetadataUrlDecision::address_reserved; + } + if (octet1 == 198 && (octet2 == 18 || octet2 == 19)) { + return MetadataUrlDecision::address_reserved; + } + if (octet1 == 198 && octet2 == 51 && octet3 == 100) { + return MetadataUrlDecision::address_reserved; + } + if (octet1 == 203 && octet2 == 0 && octet3 == 113) { + return MetadataUrlDecision::address_reserved; + } + return MetadataUrlDecision::allowed; +} + +/// Recognize the ::ffff:a.b.c.d form so an IPv4 target cannot be smuggled through an IPv6 literal. +bool mapped_v4(const net::ip::address_v6& address, net::ip::address_v4& mapped) { + const auto bytes = address.to_bytes(); + for (std::size_t index = 0; index < 10; ++index) { + if (bytes[index] != 0) { + return false; + } + } + if (bytes[10] != 0xFF || bytes[11] != 0xFF) { + return false; + } + const net::ip::address_v4::bytes_type v4_bytes{bytes[12], bytes[13], bytes[14], bytes[15]}; + mapped = net::ip::address_v4(v4_bytes); + return true; +} + +MetadataUrlDecision classify_v6(const MetadataFetchPolicy& policy, const net::ip::address_v6& address) { + net::ip::address_v4 mapped; + if (mapped_v4(address, mapped)) { + return classify_v4(policy, mapped); + } + if (address.is_unspecified()) { + return MetadataUrlDecision::address_reserved; + } + if (address.is_loopback()) { + return policy.allow_plain_http_loopback ? MetadataUrlDecision::allowed + : MetadataUrlDecision::address_loopback; + } + if (address.is_link_local()) { + return MetadataUrlDecision::address_link_local; + } + if (address.is_multicast()) { + return MetadataUrlDecision::address_multicast; + } + const auto bytes = address.to_bytes(); + if ((bytes[0] & 0xFEU) == 0xFCU) { + return MetadataUrlDecision::address_private; + } + if (address.is_site_local()) { + return MetadataUrlDecision::address_private; + } + return MetadataUrlDecision::allowed; +} + +bool is_loopback_host(const std::string& host) { + if (host == "localhost") { + return true; + } + boost::system::error_code error; + const auto address = net::ip::make_address(host, error); + return !error && address.is_loopback(); +} + +} // namespace + +std::string_view describe(MetadataUrlDecision decision) { + switch (decision) { + case MetadataUrlDecision::allowed: + return "allowed"; + case MetadataUrlDecision::malformed_url: + return "URL is malformed or carries userinfo"; + case MetadataUrlDecision::scheme_not_allowed: + return "scheme is not https and the loopback opt-out does not apply"; + case MetadataUrlDecision::origin_denied: + return "origin is on the application deny list"; + case MetadataUrlDecision::origin_not_allowed: + return "origin is not on the application allow list"; + case MetadataUrlDecision::address_link_local: + return "address is link-local"; + case MetadataUrlDecision::address_private: + return "address is in a private range"; + case MetadataUrlDecision::address_loopback: + return "address is loopback and the loopback opt-out is not enabled"; + case MetadataUrlDecision::address_multicast: + return "address is multicast or broadcast"; + case MetadataUrlDecision::address_reserved: + return "address is reserved or non-routable"; + case MetadataUrlDecision::redirect_limit_exceeded: + return "redirect chain exceeded the configured bound"; + case MetadataUrlDecision::response_too_large: + return "response exceeded the configured size cap"; + case MetadataUrlDecision::denied_origin_entry_malformed: + return "deny list entry is not a bare origin"; + } + return "refused"; +} + +std::string metadata_url_origin(const std::string& url) { + std::string scheme; + std::string authority; + if (!split_url(url, scheme, authority)) { + return {}; + } + return scheme + "://" + authority; +} + +MetadataUrlDecision validate_metadata_url(const MetadataFetchPolicy& policy, const std::string& url) { + std::string scheme; + std::string authority; + if (!split_url(url, scheme, authority)) { + return MetadataUrlDecision::malformed_url; + } + // Userinfo is the classic way to make an allowed origin read as the prefix of a hostile one. + if (authority.find('@') != std::string::npos) { + return MetadataUrlDecision::malformed_url; + } + + std::string host; + if (!authority_host(authority, host)) { + return MetadataUrlDecision::malformed_url; + } + + // Canonicalized once, before any list or callback sees it, so a deny entry cannot be bypassed by + // a differently-cased scheme/host, an explicit or oddly-spelled default port, a run of trailing + // FQDN dots, or an alternate textual spelling of the same IP literal. A URL that only canonicalizes + // this far because of a malformed port or an all-dots host is refused outright. + const auto canonical = canonicalize_origin_parts(scheme + "://" + authority); + if (!canonical) { + return MetadataUrlDecision::malformed_url; + } + const auto canonical_origin = format_canonical_origin(*canonical); + if (denies_origin(policy.denied_origins, canonical_origin)) { + return MetadataUrlDecision::origin_denied; + } + if (!contains_origin(policy.allowed_origins, canonical_origin) && + !(policy.origin_allowance && policy.origin_allowance(canonical_origin))) { + return MetadataUrlDecision::origin_not_allowed; + } + + // Runs on the canonical scheme/host, matching what the origin decision and the origin_allowance + // callback above just saw, rather than re-deriving the answer from raw, non-canonical text. + if (canonical->scheme != "https") { + const auto loopback_opt_out = policy.allow_plain_http_loopback && canonical->scheme == "http" && + is_loopback_host(canonical->host); + if (!loopback_opt_out) { + return MetadataUrlDecision::scheme_not_allowed; + } + } + + // A host written as an IP literal needs no lookup, so classify it here and refuse before + // resolution is ever reached. + boost::system::error_code error; + const auto address = net::ip::make_address(host, error); + if (!error) { + return validate_metadata_address(policy, address.to_string()); + } + return MetadataUrlDecision::allowed; +} + +MetadataUrlDecision validate_metadata_address(const MetadataFetchPolicy& policy, + const std::string& address_literal) { + boost::system::error_code error; + const auto address = net::ip::make_address(address_literal, error); + if (error) { + return MetadataUrlDecision::malformed_url; + } + if (address.is_v6()) { + return classify_v6(policy, address.to_v6()); + } + return classify_v4(policy, address.to_v4()); +} + +namespace detail { + +namespace { + +/// One decoded UTF-8 sequence. `length` is 0 when the bytes at the offset are not a well-formed +/// sequence, in which case `codepoint` is meaningless. +struct Utf8Sequence { + std::size_t length{0}; + std::uint32_t codepoint{0}; +}; + +/// Decode the UTF-8 sequence starting at `index`, rejecting every ill-formed encoding rather than +/// accepting the bytes and hoping: a truncated tail, a bad continuation byte, an overlong form, a +/// surrogate, or a value past U+10FFFF. Ill-formed input is what a peer sends when it wants the +/// consumer of the log, not the log line itself, to misbehave. +Utf8Sequence decode_utf8(std::string_view value, std::size_t index) { + const auto lead = static_cast(value[index]); + if (lead < 0x80) { + return {1, lead}; + } + + std::size_t length = 0; + std::uint32_t codepoint = 0; + if ((lead & 0xE0) == 0xC0) { + length = 2; + codepoint = lead & 0x1FU; + } else if ((lead & 0xF0) == 0xE0) { + length = 3; + codepoint = lead & 0x0FU; + } else if ((lead & 0xF8) == 0xF0) { + length = 4; + codepoint = lead & 0x07U; + } else { + return {}; // A continuation byte with no lead, or an invalid lead. + } + + if (index + length > value.size()) { + return {}; + } + for (std::size_t offset = 1; offset < length; ++offset) { + const auto continuation = static_cast(value[index + offset]); + if ((continuation & 0xC0) != 0x80) { + return {}; + } + codepoint = (codepoint << 6U) | (continuation & 0x3FU); + } + + const bool overlong = (length == 2 && codepoint < 0x80) || (length == 3 && codepoint < 0x800) || + (length == 4 && codepoint < 0x10000); + const bool surrogate = codepoint >= 0xD800 && codepoint <= 0xDFFF; + if (overlong || surrogate || codepoint > 0x10FFFF) { + return {}; + } + return {length, codepoint}; +} + +/// Codepoints that can end a line somewhere downstream. C0 and DEL are the obvious ones; C1 and +/// U+2028/U+2029 are here because a JSON encoder emits them literally and JavaScript-based log +/// viewers treat the last two as line terminators, which is the same forgery by another route. +bool breaks_a_line(std::uint32_t codepoint) { + return codepoint < 0x20 || codepoint == 0x7f || (codepoint >= 0x80 && codepoint <= 0x9f) || + codepoint == 0x2028 || codepoint == 0x2029; +} + +/// Codepoints that re-order the text after them without ending the line: the explicit embeddings and +/// overrides (U+202A-U+202E), the isolates (U+2066-U+2069) and the implicit marks LRM (U+200E), RLM +/// (U+200F) and ALM (U+061C). They let a peer reverse the rendering of the rest of a diagnostic, so a +/// refusal can be made to read as its opposite. ALM sits in the Arabic block but is the same class as +/// RLM. +bool reorders_the_line(std::uint32_t codepoint) { + return (codepoint >= 0x202a && codepoint <= 0x202e) || + (codepoint >= 0x2066 && codepoint <= 0x2069) || codepoint == 0x200e || codepoint == 0x200f || + codepoint == 0x061c; +} + +} // namespace + +std::string sanitize_for_diagnostics(std::string_view value) { + constexpr std::size_t max_length = 256; + std::string cleaned; + cleaned.reserve(std::min(value.size(), max_length)); + + std::size_t index = 0; + bool truncated = false; + while (index < value.size()) { + const auto decoded = decode_utf8(value, index); + + // Ill-formed bytes are replaced one for one rather than copied, so the result is always + // well-formed UTF-8 even when the input was not. Text that reaches here through a header + // rather than through a JSON document has had nothing validate it. + const std::size_t width = decoded.length == 0 ? 1 : decoded.length; + if (cleaned.size() + width > max_length) { + truncated = true; + break; + } + if (decoded.length == 0) { + cleaned.push_back('?'); + } else if (breaks_a_line(decoded.codepoint) || reorders_the_line(decoded.codepoint)) { + cleaned.push_back(' '); + } else { + cleaned.append(value.substr(index, decoded.length)); + } + index += width; + } + + if (truncated) { + cleaned += "..."; + } + return cleaned; +} + +} // namespace detail + +/// The target is sanitized for the MESSAGE and kept raw in `target_`. Every throw site passes a URL +/// or address that came from a peer, and sanitizing here rather than at each of them means a throw +/// site cannot forget. `target()` still returns the value verbatim, because a caller inspecting it +/// programmatically wants the real URL, not one with its control characters flattened. +MetadataPolicyError::MetadataPolicyError(MetadataUrlDecision decision, std::string target) + : std::runtime_error("OAuth metadata target refused (" + std::string(describe(decision)) + + "): " + detail::sanitize_for_diagnostics(target)), + decision_(decision), + target_(std::move(target)) {} + +} // namespace mcp::auth diff --git a/src/auth/oauth.cpp b/src/auth/oauth.cpp new file mode 100644 index 0000000..6761002 --- /dev/null +++ b/src/auth/oauth.cpp @@ -0,0 +1,2562 @@ +#include + +#include "oauth_internal.hpp" + +#include + +#include +#include + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +// GCC 11 SSO Coroutine Safety -- see docs/contributing.rst "Known Issues". +// Strings and string-containing protocol values that cross a suspension point +// live in shared operation state rather than directly in coroutine frames. + +namespace mcp::auth { + +namespace beast = boost::beast; +namespace http = beast::http; +namespace net = boost::asio; + +namespace detail { + +std::string base64_encode(const unsigned char* data, std::size_t len) { + std::string result; + result.reserve(((len + 2) / 3) * 4); + + for (std::size_t i = 0; i < len; i += 3) { + unsigned int value = static_cast(data[i]) << constants::g_shift16; + if (i + 1 < len) { + value |= static_cast(data[i + 1]) << constants::g_shift8; + } + if (i + 2 < len) { + value |= static_cast(data[i + 2]); + } + + result.push_back( + mcp::constants::g_alphabet[(value >> constants::g_shift18) & constants::g_mask0x3F]); + result.push_back( + mcp::constants::g_alphabet[(value >> constants::g_shift12) & constants::g_mask0x3F]); + result.push_back( + (i + 1 < len) + ? mcp::constants::g_alphabet[(value >> constants::g_shift6) & constants::g_mask0x3F] + : '='); + result.push_back((i + 2 < len) ? mcp::constants::g_alphabet[value & constants::g_mask0x3F] + : '='); + } + + return result; +} + +std::string base64url_encode(const unsigned char* data, std::size_t len) { + auto encoded = base64_encode(data, len); + + for (auto& ch : encoded) { + if (ch == '+') { + ch = '-'; + } else if (ch == '/') { + ch = '_'; + } + } + encoded.erase(std::remove(encoded.begin(), encoded.end(), '='), encoded.end()); + return encoded; +} + +std::array sha256(const std::string& input) { + std::array digest{}; + unsigned int digest_len = 0; + + std::unique_ptr context(EVP_MD_CTX_new(), EVP_MD_CTX_free); + if (!context || EVP_DigestInit_ex(context.get(), EVP_sha256(), nullptr) != 1 || + EVP_DigestUpdate(context.get(), input.data(), input.size()) != 1 || + EVP_DigestFinal_ex(context.get(), digest.data(), &digest_len) != 1) { + throw std::runtime_error("OpenSSL SHA-256 failed"); + } + + return digest; +} + +std::string generate_random_string(std::size_t length) { + static const std::size_t s_charset_size = mcp::constants::g_unreserved_chars.size(); + static const auto s_bias_limit = + static_cast((256 / s_charset_size) * s_charset_size); + + std::string result; + result.reserve(length); + while (result.size() < length) { + unsigned char byte = 0; + if (RAND_bytes(&byte, 1) != 1) { + throw std::runtime_error("RAND_bytes failed"); + } + if (byte < s_bias_limit) { + result.push_back(mcp::constants::g_unreserved_chars[byte % s_charset_size]); + } + } + return result; +} + +std::string url_encode(const std::string& value) { + std::string result; + result.reserve(value.size() * 3); + + for (unsigned char ch : value) { + if ((ch >= 'A' && ch <= 'Z') || (ch >= 'a' && ch <= 'z') || (ch >= '0' && ch <= '9') || + ch == '-' || ch == '_' || ch == '.' || ch == '~') { + result.push_back(static_cast(ch)); + } else { + result.push_back('%'); + result.push_back(mcp::constants::g_hex_digits_upper[ch >> constants::g_shift4]); + result.push_back(mcp::constants::g_hex_digits_upper[ch & constants::g_mask0x0F]); + } + } + return result; +} + +std::string build_form_body(const KeyValuePairList& params) { + std::string body; + for (const auto& [key, value] : params) { + if (!body.empty()) { + body.push_back('&'); + } + body += url_encode(key) + "=" + url_encode(value); + } + return body; +} + +} // namespace detail + +PkcePair generate_pkce_pair(std::size_t verifier_length) { + if (verifier_length < constants::g_min_verifier_length || + verifier_length > constants::g_max_verifier_length) { + throw std::invalid_argument("PKCE verifier length must be 43-128 characters"); + } + + PkcePair pair; + pair.code_verifier = detail::generate_random_string(verifier_length); + pair.challenge_method = "S256"; + + const auto hash = detail::sha256(pair.code_verifier); + pair.code_challenge = detail::base64url_encode(hash.data(), hash.size()); + return pair; +} + +bool TokenResponse::is_expired(int margin) const { + if (!expires_in.has_value()) { + return false; + } + const auto expiry = received_at + std::chrono::seconds(*expires_in) - std::chrono::seconds(margin); + return std::chrono::steady_clock::now() >= expiry; +} + +void from_json(const nlohmann::json& json, TokenResponse& token) { + json.at("access_token").get_to(token.access_token); + token.token_type = json.value("token_type", "Bearer"); + if (json.contains("refresh_token")) { + token.refresh_token = json.at("refresh_token").get(); + } + if (json.contains("expires_in")) { + token.expires_in = json.at("expires_in").get(); + } + if (json.contains("scope")) { + token.scope = json.at("scope").get(); + } + token.received_at = std::chrono::steady_clock::now(); +} + +void to_json(nlohmann::json& json, const TokenResponse& token) { + json = nlohmann::json{{"access_token", token.access_token}, {"token_type", token.token_type}}; + if (token.refresh_token) { + json["refresh_token"] = *token.refresh_token; + } + if (token.expires_in) { + json["expires_in"] = *token.expires_in; + } + if (token.scope) { + json["scope"] = *token.scope; + } +} + +struct InMemoryTokenStore::Impl { + mutable std::mutex mutex; + std::unordered_map tokens; +}; + +InMemoryTokenStore::InMemoryTokenStore() : impl_(std::make_unique()) {} + +InMemoryTokenStore::~InMemoryTokenStore() = default; + +void InMemoryTokenStore::store(const std::string& server_url, TokenResponse token) { + std::lock_guard lock(impl_->mutex); + impl_->tokens[server_url] = std::move(token); +} + +std::optional InMemoryTokenStore::load(const std::string& server_url) const { + std::lock_guard lock(impl_->mutex); + const auto iter = impl_->tokens.find(server_url); + if (iter == impl_->tokens.end()) { + return std::nullopt; + } + return iter->second; +} + +void InMemoryTokenStore::remove(const std::string& server_url) { + std::lock_guard lock(impl_->mutex); + impl_->tokens.erase(server_url); +} + +namespace { + +constexpr std::string_view g_https_prefix = "https://"; + +/// Flattens peer-controlled text -- control characters, bidi overrides and ill-formed UTF-8 -- and +/// caps it, so a value a peer chose cannot forge or reorder a line of a diagnostic. Every throw +/// site in this file that interpolates a value originating from a peer must route it through here; +/// `MetadataPolicyError` does the same inside its own constructor, so its throw sites do not repeat +/// it. This is a convention, not something the compiler checks: a new throw site that forgets it +/// re-opens the hole silently. +using detail::sanitize_for_diagnostics; + +/// Resolve a `Location` header against the URL that produced it. +std::string resolve_redirect_target(const std::string& base, const std::string& location) { + if (location.find("://") != std::string::npos) { + return location; + } + const auto origin = metadata_url_origin(base); + if (origin.empty()) { + return location; + } + if (!location.empty() && location.front() == '/') { + return origin + location; + } + + auto directory = base; + if (const auto query = directory.find_first_of("?#"); query != std::string::npos) { + directory.erase(query); + } + const auto last_slash = directory.rfind('/'); + if (last_slash == std::string::npos || last_slash < origin.size()) { + return origin + "/" + location; + } + return directory.substr(0, last_slash + 1) + location; +} + +bool is_redirect_status(unsigned int status) { + return status == static_cast(http::status::moved_permanently) || + status == static_cast(http::status::found) || + status == static_cast(http::status::see_other) || + status == static_cast(http::status::temporary_redirect) || + status == static_cast(http::status::permanent_redirect); +} + +/// RFC 9728 section 3.3: does a protected-resource metadata `resource` value identify the server this +/// client is configured for? Accepted when the two are byte-exact, or when `resource` is a proper URI +/// prefix of `server_url` under RFC 8707 audience semantics: same scheme and authority, compared +/// byte-exact with no normalization (an explicit default port differs from an implicit one), and a +/// path that is empty/`"/"` (matches any path) or a prefix of the server URL's path on a `/` segment +/// boundary. A single trailing `/` on `resource`'s path is ignored. `resource` may not carry a query +/// or fragment. +/// +/// An origin-root resource is accepted for any path by design: the root-PRM layout (conformance +/// `auth/metadata-var2`). This only decides whether the value may be sent as the `resource` +/// parameter; the authorization server remains the arbiter of the audience it issues. +bool resource_identifies_server(const std::string& resource, const std::string& server_url) { + if (resource == server_url) { + return true; + } + if (resource.find('?') != std::string::npos || resource.find('#') != std::string::npos) { + return false; + } + const auto resource_origin = metadata_url_origin(resource); + const auto server_origin = metadata_url_origin(server_url); + if (resource_origin.empty() || resource_origin != server_origin) { + return false; + } + auto resource_path = resource.substr(resource_origin.size()); + if (resource_path.empty() || resource_path == "/") { + return true; + } + if (resource_path.back() == '/') { + resource_path.pop_back(); + } + const auto server_path = server_url.substr(server_origin.size()); + if (server_path == resource_path) { + return true; + } + return server_path.size() > resource_path.size() && + server_path.compare(0, resource_path.size(), resource_path) == 0 && + server_path[resource_path.size()] == '/'; +} + +/// Throws unless `resource` is present and identifies `server_url`. +/// +/// Called before the SDK contacts anything the protected-resource document names, so a document +/// that does not identify our configured server cannot make us fetch attacker-named authorization +/// server metadata, register a client with it, or persist those credentials. +void require_resource_identifies_server(const std::string& resource, const std::string& server_url) { + if (resource.empty()) { + // RFC 9728 §2 makes `resource` a required member; a PRM that omits it (or ships it + // empty) does not meet the spec and must not be trusted silently. + throw std::runtime_error("Protected resource metadata for server '" + + sanitize_for_diagnostics(server_url) + + "' is missing the required 'resource' member"); + } + if (!resource_identifies_server(resource, server_url)) { + throw std::runtime_error("Protected resource metadata resource '" + + sanitize_for_diagnostics(resource) + "' does not identify server '" + + sanitize_for_diagnostics(server_url) + "'"); + } +} + +} // namespace + +namespace detail { + +/// One scope's abort latch, owned jointly by the scope object and by every exchange issued +/// through it. `aborted` is guarded by the owning client's `active_mutex`, the same mutex the +/// client-wide latch is read and written under, so the scoped and unscoped checks keep the +/// single synchronisation discipline they have always had. +struct OAuthScopeState { + OAuthScopeState() { live().fetch_add(1, std::memory_order_relaxed); } + OAuthScopeState(const OAuthScopeState&) = delete; + OAuthScopeState& operator=(const OAuthScopeState&) = delete; + ~OAuthScopeState() { live().fetch_sub(1, std::memory_order_relaxed); } + + bool aborted{false}; + + /// How many latches exist right now, across every client in the process. Instrumentation for + /// the retention test: the count returning to its starting value after a churn of scopes is + /// what says the latches were released rather than parked somewhere. + static std::atomic& live() { + static std::atomic count{0}; + return count; + } + static std::size_t live_count() { return live().load(std::memory_order_relaxed); } +}; + +} // namespace detail + +struct OAuthHttpClient::Impl : std::enable_shared_from_this { + struct ParsedUrl { + std::string scheme; + std::string host; + std::string port; + std::string path; + }; + + /// State for one HTTP exchange. `reset` rebuilds the per-hop pieces so a redirect can be + /// followed on a fresh connection without discarding the operation. + struct Exchange { + Exchange(std::shared_ptr state, net::strand executor, + std::string target, std::shared_ptr abort_scope, + MetadataFetchPolicy fetch_policy, HostResolver resolver_hook) + : owner(std::move(state)), + strand(std::move(executor)), + url(std::move(target)), + resolver(strand), + scope(std::move(abort_scope)), + policy(std::move(fetch_policy)), + host_resolver(std::move(resolver_hook)) {} + + void reset() { + endpoints.clear(); + // A redirect is followed on a fresh connection; tcp_stream is not reassignable. + stream.emplace(strand); + buffer.clear(); + parser.emplace(); + parser->body_limit(policy.max_response_bytes); + } + + [[nodiscard]] const http::response& response() const { + return parser->get(); + } + + std::shared_ptr owner; + net::strand strand; + ParsedUrl parsed; + std::string url; + net::ip::tcp::resolver resolver; + std::vector endpoints; + std::optional stream; + beast::flat_buffer buffer; + std::optional> parser; + std::string body; + /// The abort latch this exchange answers to, held for the exchange's whole life so the + /// latch cannot be freed while this exchange can still read it. Null is the client's own + /// unscoped work, which only abort_pending() ends. + std::shared_ptr scope; + /// Pinned for this exchange's whole life; see Impl::make_exchange(). + MetadataFetchPolicy policy; + HostResolver host_resolver; + }; + + explicit Impl(const net::any_io_executor& executor) : strand(net::make_strand(executor)) {} + + static ParsedUrl parse_url(const std::string& url) { + std::string scheme; + std::string default_port; + std::size_t prefix_length = 0; + if (url.starts_with(mcp::constants::g_http_prefix)) { + scheme = "http"; + default_port = "80"; + prefix_length = mcp::constants::g_http_prefix.size(); + } else if (url.starts_with(g_https_prefix)) { + scheme = "https"; + default_port = "443"; + prefix_length = g_https_prefix.size(); + } else { + throw std::invalid_argument( + "OAuth HTTP client URL must use the http:// or https:// scheme"); + } + + auto authority_and_path = url.substr(prefix_length); + const auto path_separator = authority_and_path.find('/'); + auto authority = authority_and_path.substr(0, path_separator); + auto path = + path_separator == std::string::npos ? "/" : authority_and_path.substr(path_separator); + + std::string host; + auto port = std::move(default_port); + const auto colon = authority.find(':'); + if (colon == std::string::npos) { + host = std::move(authority); + } else { + host = authority.substr(0, colon); + port = authority.substr(colon + 1); + } + + return {std::move(scheme), std::move(host), std::move(port), std::move(path)}; + } + + /// Validate a target before any lookup. Refusal happens here, so the host of a refused target + /// is never resolved and no socket is opened for it. + static void enforce_url_policy(const MetadataFetchPolicy& policy, const std::string& url) { + const auto decision = validate_metadata_url(policy, url); + if (decision != MetadataUrlDecision::allowed) { + throw MetadataPolicyError(decision, url); + } + } + + /// Classify every address a lookup produced. One blocked answer refuses the whole fetch, so a + /// resolver that mixes a routable answer with a hostile one cannot smuggle the hostile one in. + static void enforce_address_policy(const MetadataFetchPolicy& policy, + const std::vector& endpoints) { + for (const auto& endpoint : endpoints) { + auto literal = endpoint.address().to_string(); + const auto decision = validate_metadata_address(policy, literal); + if (decision != MetadataUrlDecision::allowed) { + throw MetadataPolicyError(decision, std::move(literal)); + } + } + } + + static Task connect(const std::shared_ptr& exchange) { + if (exchange->parsed.scheme != "http") { + throw std::runtime_error( + "OAuth over https requires TLS support, which this build does not provide: " + + sanitize_for_diagnostics(exchange->url)); + } + + if (exchange->host_resolver) { + const auto port = static_cast(std::stoul(exchange->parsed.port)); + for (const auto& literal : + exchange->host_resolver(exchange->parsed.host, exchange->parsed.port)) { + boost::system::error_code parse_error; + const auto address = net::ip::make_address(literal, parse_error); + if (parse_error) { + throw std::runtime_error("Host resolver returned an unusable address: " + + sanitize_for_diagnostics(literal)); + } + exchange->endpoints.emplace_back(address, port); + } + } else { + const auto results = co_await exchange->resolver.async_resolve( + exchange->parsed.host, exchange->parsed.port, net::use_awaitable); + for (const auto& entry : results) { + exchange->endpoints.push_back(entry.endpoint()); + } + } + + if (exchange->endpoints.empty()) { + throw std::runtime_error("No address resolved for " + + sanitize_for_diagnostics(exchange->parsed.host)); + } + enforce_address_policy(exchange->policy, exchange->endpoints); + + // Last look at the abort before a socket exists. An abort_pending() that ran while the + // lookup was past cancelling found nothing to close. This coroutine holds the client's + // strand from here until it suspends inside async_connect(), which opens the socket + // first, and the close a later abort posts runs on that strand: it cannot run before the + // socket is open, whichever thread it runs on, so it finds the socket and closes it. + throw_if_aborted(exchange); + + // Connect only to the addresses this single lookup produced. They are pinned for the + // exchange, so a name that resolves differently later cannot redirect it. + exchange->stream->expires_after(std::chrono::seconds(mcp::constants::g_http_timeout_seconds)); + co_await exchange->stream->async_connect(exchange->endpoints, net::use_awaitable); + } + + template + static Task exchange_once(const std::shared_ptr& exchange, Request& request) { + co_await connect(exchange); + + exchange->stream->expires_after(std::chrono::seconds(mcp::constants::g_http_timeout_seconds)); + co_await http::async_write(*exchange->stream, request, net::use_awaitable); + try { + co_await http::async_read(*exchange->stream, exchange->buffer, *exchange->parser, + net::use_awaitable); + } catch (const boost::system::system_error& error) { + beast::error_code discard_error; + (void)exchange->stream->socket().shutdown(net::ip::tcp::socket::shutdown_both, + discard_error); + if (error.code() == http::error::body_limit) { + // The body is abandoned mid-read rather than buffered to completion. + throw MetadataPolicyError(MetadataUrlDecision::response_too_large, exchange->url); + } + throw; + } + + beast::error_code shutdown_error; + (void)exchange->stream->socket().shutdown(net::ip::tcp::socket::shutdown_both, shutdown_error); + } + + /// RAII membership in `active_exchanges` for the lifetime of one HTTP exchange, so + /// `abort_pending()` can reach a stalled connect/write/read and never targets one that finished. + struct ActiveExchangeGuard { + ActiveExchangeGuard(std::shared_ptr owner_in, Exchange* exchange_in) + : owner(std::move(owner_in)), exchange(exchange_in) {} + ActiveExchangeGuard(const ActiveExchangeGuard&) = delete; + ActiveExchangeGuard& operator=(const ActiveExchangeGuard&) = delete; + ActiveExchangeGuard& operator=(ActiveExchangeGuard&&) = delete; + + ActiveExchangeGuard(ActiveExchangeGuard&& other) noexcept + : owner(std::move(other.owner)), exchange(other.exchange) { + other.exchange = nullptr; + } + + ~ActiveExchangeGuard() { + if (!owner || exchange == nullptr) { + return; + } + std::lock_guard lock(owner->active_mutex); + auto& list = owner->active_exchanges; + list.erase(std::remove_if(list.begin(), list.end(), + [this](const std::weak_ptr& weak) { + const auto locked = weak.lock(); + return !locked || locked.get() == exchange; + }), + list.end()); + } + + std::shared_ptr owner; + Exchange* exchange; + }; + + /// Registers `exchange` as in flight, unless `abort_pending()` has already run -- in which case + /// this exchange is refused before it opens a connection, the same as one abort_pending() closes + /// out from under. Without this check, an exchange that starts registering after abort_pending() + /// already swept `active_exchanges` would never be reached by it and would run to completion + /// (see the sticky `aborted` flag on abort_pending() below). + static ActiveExchangeGuard track_exchange(const std::shared_ptr& exchange) { + std::lock_guard lock(exchange->owner->active_mutex); + if (exchange->owner->is_aborted(exchange->scope)) { + // The same error an exchange already in flight sees when the abort closes its socket + // underneath it, so a caller cannot tell whether this exchange started before or after + // the client -- or its own scope -- was closed. + throw boost::system::system_error(net::error::operation_aborted); + } + exchange->owner->active_exchanges.push_back(exchange); + return ActiveExchangeGuard(exchange->owner, exchange.get()); + } + + /// Whether `scope` may still issue requests. The client-wide latch ends everything; a scoped + /// latch ends only that scope. Callers hold active_mutex. + [[nodiscard]] static bool is_aborted_locked(bool client_aborted, + const std::shared_ptr& scope) { + return client_aborted || (scope && scope->aborted); + } + + [[nodiscard]] bool is_aborted(const std::shared_ptr& scope) const { + return is_aborted_locked(aborted, scope); + } + + /// Re-check the sticky abort part-way through an exchange, throwing exactly what + /// track_exchange() throws. + /// + /// track_exchange() runs once per exchange, but run_get() follows redirects in a loop with a + /// fresh socket per iteration. An abort that lands as an iteration's read completes closes a + /// finished socket, so without this check the coroutine would follow the redirect and run a new + /// request for a torn-down client. + static void throw_if_aborted(const std::shared_ptr& exchange) { + std::lock_guard lock(exchange->owner->active_mutex); + if (exchange->owner->is_aborted(exchange->scope)) { + throw boost::system::system_error(net::error::operation_aborted); + } + } + + /// Close the underlying socket of every exchange currently in flight, posted onto the client's + /// strand so the closure is never raced with the coroutine using it. A pending resolve, connect, + /// write, or read then completes with an error instead of hanging. Also latches `aborted`, so + /// every exchange that registers with track_exchange() from this point on -- including one that + /// has not made its first network call yet -- is refused rather than left to run past a client + /// that has moved on. Idempotent; a no-op when nothing is in flight either way. + void abort_pending() { abort_matching(nullptr); } + + /// The scoped counterpart, with the same guarantees confined to one scope: `scope` is latched + /// and every exchange belonging to it is closed, while other scopes and the client's own + /// unscoped work carry on. Pass `nullptr` to mean the whole client, which is what + /// abort_pending() does. + void abort_matching(const std::shared_ptr& scope) { + std::vector> exchanges; + { + std::lock_guard lock(active_mutex); + if (!scope) { + aborted = true; + } else { + scope->aborted = true; + } + for (auto& weak : active_exchanges) { + auto locked = weak.lock(); + if (locked && (!scope || locked->scope == scope)) { + exchanges.push_back(std::move(locked)); + } + } + } + for (auto& exchange : exchanges) { + net::post(strand, [exchange]() { + exchange->resolver.cancel(); + if (exchange->stream) { + beast::error_code ec; + exchange->stream->socket().close(ec); + } + }); + } + } + + /// Hand out a fresh latch. A scope is identified by the control block itself, which every party + /// that can consult it keeps alive, so a new scope can never alias a latched one the way a + /// recycled id could. Do not reintroduce scope ids. + static std::shared_ptr new_scope() { + return std::make_shared(); + } + + /// Build an exchange with the fetch policy and host resolver pinned for its whole lifetime, read + /// once under the mutex that guards the exchange list: the setters are plain writes to state the + /// exchange coroutines read, and an exchange re-validates every redirect hop, so a policy swapped + /// mid-chain would otherwise check later hops against different rules. Every request carries the + /// scope it was issued under, so aborting that scope reaches exactly these exchanges; a null + /// scope is the client's own unscoped work. + std::shared_ptr make_exchange(std::string url, + std::shared_ptr scope) { + std::lock_guard lock(active_mutex); + return std::make_shared(shared_from_this(), strand, std::move(url), std::move(scope), + policy, host_resolver); + } + + /// Run one exchange on the client's strand and hand its outcome back to the caller off it. + /// + /// The strand keeps an exchange apart from the close abort_matching() posts when several threads + /// run the io_context. Posting the coroutine to the strand is not enough: an awaitable's executor + /// is fixed when it is spawned, so it is spawned on the strand and every resumption lands there. + /// + /// co_spawn() completes by dispatching to the caller's executor, which on an io_context thread + /// runs the caller in place inside the strand handler; the caller is posted off the strand first + /// so it does not hold it. The result waits out that post on the heap, not in this frame; see the + /// GCC 11 note at the top of this file. + template + static Task on_strand(net::strand strand, Task exchange) { + std::shared_ptr result; + std::exception_ptr failure; + try { + result = std::make_shared( + co_await net::co_spawn(strand, std::move(exchange), net::use_awaitable)); + } catch (...) { + failure = std::current_exception(); + } + if (strand.running_in_this_thread()) { +#if BOOST_VERSION >= 107700 + // A caller that has been cancelled still has to leave the strand, and co_await throws + // for a cancelled coroutine before it initiates anything. The outcome is already + // decided either way: the cancellation reached the exchange through co_spawn(), or + // arrived too late to matter to it. + // + // The setting belongs to the caller's whole coroutine, so it goes back as it was + // found even if the post throws. Restoring it is itself a co_await, which rules out + // a destructor or a catch block. + const bool throws_if_cancelled = co_await net::this_coro::throw_if_cancelled(); + co_await net::this_coro::throw_if_cancelled(false); + std::exception_ptr hop_failure; + try { + co_await net::post(net::use_awaitable); + } catch (...) { + hop_failure = std::current_exception(); + } + co_await net::this_coro::throw_if_cancelled(throws_if_cancelled); + if (hop_failure) { + std::rethrow_exception(hop_failure); + } +#else + co_await net::post(net::use_awaitable); +#endif + } + if (failure) { + std::rethrow_exception(failure); + } + co_return std::move(*result); + } + + Task get_json(std::string url, + std::shared_ptr scope = nullptr) { + return on_strand(strand, run_get(make_exchange(std::move(url), std::move(scope)))); + } + + Task post_token_request(std::string token_endpoint, const KeyValuePairList& params, + std::string authorization, + std::shared_ptr scope = nullptr) { + return on_strand(strand, + run_post(make_exchange(std::move(token_endpoint), std::move(scope)), + std::make_shared(detail::build_form_body(params)), + std::move(authorization))); + } + + Task post_json(std::string url, std::string body, + std::shared_ptr scope = nullptr) { + // Sanitized where the label is built, not where it is thrown: `url` is moved into the + // exchange on the next line, and the throw site only ever sees the label. + auto label = "HTTP POST " + sanitize_for_diagnostics(url); + return on_strand(strand, run_post_json(make_exchange(std::move(url), std::move(scope)), + std::make_shared(std::move(body)), + "application/json", std::move(label))); + } + + static Task run_get(std::shared_ptr exchange) { + auto guard = track_exchange(exchange); + + const auto redirect_budget = exchange->policy.max_redirects; + for (std::size_t redirect = 0;; ++redirect) { + // Re-read the abort on every hop, not just at track_exchange() above: abort_pending() + // can land between two iterations, where it has nothing left to close. + throw_if_aborted(exchange); + // Every hop, including each redirect target, is validated afresh before it is reached. + enforce_url_policy(exchange->policy, exchange->url); + exchange->parsed = parse_url(exchange->url); + exchange->reset(); + + http::request request(http::verb::get, exchange->parsed.path, + mcp::constants::g_http_version_11); + request.set(http::field::host, exchange->parsed.host); + request.set(http::field::accept, "application/json"); + co_await exchange_once(exchange, request); + + const auto status = exchange->response().result_int(); + if (!is_redirect_status(status)) { + break; + } + if (redirect >= redirect_budget) { + throw MetadataPolicyError(MetadataUrlDecision::redirect_limit_exceeded, exchange->url); + } + const auto location = exchange->response().find(http::field::location); + if (location == exchange->response().end()) { + throw std::runtime_error("HTTP GET " + sanitize_for_diagnostics(exchange->url) + + " returned a redirect without a Location header"); + } + exchange->url = resolve_redirect_target(exchange->url, std::string(location->value())); + } + + if (exchange->response().result_int() >= mcp::constants::g_http_bad_request) { + throw std::runtime_error("HTTP GET " + sanitize_for_diagnostics(exchange->url) + + " failed with status " + + std::to_string(exchange->response().result_int())); + } + + exchange->body = exchange->response().body(); + auto response_json = nlohmann::json::parse(exchange->body, nullptr, false); + if (response_json.is_discarded()) { + throw std::runtime_error("Failed to parse JSON from " + + sanitize_for_diagnostics(exchange->url)); + } + co_return response_json; + } + + /// One POST exchange returning the parsed JSON reply. Redirects are deliberately not followed: + /// replaying a token or registration body at a redirect target would hand its contents to + /// whatever the first hop nominated. + static Task run_post_json(std::shared_ptr exchange, + std::shared_ptr body, + std::string content_type, std::string failure_label, + std::string authorization = {}) { + auto guard = track_exchange(exchange); + + enforce_url_policy(exchange->policy, exchange->url); + exchange->parsed = parse_url(exchange->url); + exchange->reset(); + + http::request request(http::verb::post, exchange->parsed.path, + mcp::constants::g_http_version_11); + request.set(http::field::host, exchange->parsed.host); + request.set(http::field::content_type, content_type); + request.set(http::field::accept, "application/json"); + if (!authorization.empty()) { + request.set(http::field::authorization, authorization); + } + request.body() = *body; + request.prepare_payload(); + co_await exchange_once(exchange, request); + + if (exchange->response().result_int() >= mcp::constants::g_http_bad_request) { + throw std::runtime_error(failure_label + " failed with status " + + std::to_string(exchange->response().result_int()) + ": " + + sanitize_for_diagnostics(exchange->response().body())); + } + + exchange->body = exchange->response().body(); + auto response_json = nlohmann::json::parse(exchange->body, nullptr, false); + if (response_json.is_discarded()) { + throw std::runtime_error("Failed to parse JSON from " + + sanitize_for_diagnostics(exchange->url)); + } + co_return response_json; + } + + static Task run_post(std::shared_ptr exchange, + std::shared_ptr form_body, + std::string authorization) { + auto response_json = co_await run_post_json(std::move(exchange), std::move(form_body), + "application/x-www-form-urlencoded", + "Token request", std::move(authorization)); + co_return response_json.get(); + } + + net::strand strand; + std::mutex active_mutex; + /// Both guarded by active_mutex: written by the setters, read once per exchange in + /// make_exchange(). + MetadataFetchPolicy policy; + HostResolver host_resolver; + std::vector> active_exchanges; + /// Set once by abort_pending(); makes the abort sticky so an exchange started afterward is + /// refused instead of running to completion. Guarded by active_mutex alongside the list above. + bool aborted{false}; + // A scope's abort latch lives on the scope itself (detail::OAuthScopeState). Do not add a + // container of scope state here: it would have to outlive every exchange that could still + // consult it, which in practice means never being cleaned up. +}; + +OAuthHttpClient::OAuthHttpClient(const net::any_io_executor& executor) + : impl_(std::make_shared(executor)) {} + +/// Safe to call at any time, including mid-flight: exchanges already running keep what they started +/// with, and the next one picks up the new value. See configure() below for changing both together. +void OAuthHttpClient::set_metadata_policy(MetadataFetchPolicy policy) { + std::lock_guard lock(impl_->active_mutex); + impl_->policy = std::move(policy); +} + +void OAuthHttpClient::set_host_resolver(HostResolver resolver) { + std::lock_guard lock(impl_->active_mutex); + impl_->host_resolver = std::move(resolver); +} + +/// Installs both values under one lock. Use this rather than the two setters above whenever both +/// change: calling them in sequence leaves an interval holding one new value and one old one, and +/// make_exchange() pins whatever it finds. Resolver-first is the natural order to write and the +/// dangerous one. See the @warning on set_metadata_policy() in include/mcp/auth/oauth.hpp for why, +/// and for how wide the interval measures. +void OAuthHttpClient::configure(MetadataFetchPolicy policy, HostResolver resolver) { + std::lock_guard lock(impl_->active_mutex); + impl_->policy = std::move(policy); + impl_->host_resolver = std::move(resolver); +} + +void OAuthHttpClient::abort_pending() { impl_->abort_pending(); } + +namespace { + +/// Place the client secret where the negotiated method says it belongs, and nowhere else. +/// +/// @return The `Authorization` header value, empty when the secret does not travel in a header. +std::string apply_client_authentication(const OAuthConfig& config, KeyValuePairList& params) { + const auto& method = config.token_endpoint_auth_method; + if (method && *method == "none") { + return {}; + } + if (!config.client_secret) { + return {}; + } + if (method && *method == "client_secret_basic") { + const auto credentials = + detail::url_encode(config.client_id) + ":" + detail::url_encode(*config.client_secret); + return "Basic " + + detail::base64_encode(reinterpret_cast(credentials.data()), + credentials.size()); + } + params.emplace_back("client_secret", *config.client_secret); + return {}; +} + +} // namespace + +namespace { + +/// The form body and Authorization header for one token-endpoint request. Built here rather than +/// inline so the client and a scope on it issue byte-identical requests, differing only in which +/// abort scope the exchange is tracked under. +struct TokenRequest { + KeyValuePairList params; + std::string authorization; +}; + +TokenRequest build_authorization_code_request(const OAuthConfig& config, const std::string& code, + const std::string& code_verifier) { + KeyValuePairList params = { + {"grant_type", "authorization_code"}, {"code", code}, + {"redirect_uri", config.redirect_uri}, {"client_id", config.client_id}, + {"code_verifier", code_verifier}, + }; + auto authorization = apply_client_authentication(config, params); + if (config.resource) { + params.emplace_back("resource", *config.resource); + } + return {std::move(params), std::move(authorization)}; +} + +TokenRequest build_refresh_request(const OAuthConfig& config, const std::string& refresh_token) { + KeyValuePairList params = { + {"grant_type", "refresh_token"}, + {"refresh_token", refresh_token}, + {"client_id", config.client_id}, + }; + auto authorization = apply_client_authentication(config, params); + if (config.resource) { + params.emplace_back("resource", *config.resource); + } + return {std::move(params), std::move(authorization)}; +} + +} // namespace + +Task OAuthHttpClient::exchange_code(const OAuthConfig& config, const std::string& code, + const std::string& code_verifier) { + auto request = build_authorization_code_request(config, code, code_verifier); + return impl_->post_token_request(config.token_endpoint, request.params, + std::move(request.authorization)); +} + +Task OAuthHttpClient::refresh_token(const OAuthConfig& config, + const std::string& refresh_token) { + auto request = build_refresh_request(config, refresh_token); + return impl_->post_token_request(config.token_endpoint, request.params, + std::move(request.authorization)); +} + +Task OAuthHttpClient::get_json(const std::string& url) { return impl_->get_json(url); } + +Task OAuthHttpClient::post_json(const std::string& url, const nlohmann::json& body) { + return impl_->post_json(url, body.dump()); +} + +namespace detail { + +/// Defined here and nowhere else; the public header only grants it friendship. +struct OAuthTestAccess { + static std::size_t retained_scope_records(const OAuthHttpClient& client) { + // Structural, not a stubbed zero: there is no per-scope container to count. A change + // that reintroduces one has to answer here, and the retention test will see it. + std::lock_guard lock(client.impl_->active_mutex); + return 0; + } +}; + +} // namespace detail + +namespace internal { + +std::size_t retained_scope_record_count(const OAuthHttpClient& client) { + return detail::OAuthTestAccess::retained_scope_records(client); +} + +std::size_t live_scope_latch_count() { return detail::OAuthScopeState::live_count(); } + +} // namespace internal + +OAuthHttpClientScope OAuthHttpClient::make_scope() { + return OAuthHttpClientScope(impl_, impl_->new_scope()); +} + +OAuthHttpClientScope::OAuthHttpClientScope(std::shared_ptr impl, + std::shared_ptr state) + : impl_(std::move(impl)), state_(std::move(state)) {} + +Task OAuthHttpClientScope::exchange_code(const OAuthConfig& config, + const std::string& code, + const std::string& code_verifier) { + auto request = build_authorization_code_request(config, code, code_verifier); + return impl_->post_token_request(config.token_endpoint, request.params, + std::move(request.authorization), state_); +} + +Task OAuthHttpClientScope::refresh_token(const OAuthConfig& config, + const std::string& refresh_token) { + auto request = build_refresh_request(config, refresh_token); + return impl_->post_token_request(config.token_endpoint, request.params, + std::move(request.authorization), state_); +} + +Task OAuthHttpClientScope::get_json(const std::string& url) { + return impl_->get_json(url, state_); +} + +Task OAuthHttpClientScope::post_json(const std::string& url, + const nlohmann::json& body) { + return impl_->post_json(url, body.dump(), state_); +} + +void OAuthHttpClientScope::abort() { impl_->abort_matching(state_); } + +void from_json(const nlohmann::json& json, ProtectedResourceMetadata& metadata) { + metadata.raw = json; + if (json.contains("resource")) { + json.at("resource").get_to(metadata.resource); + } + if (json.contains("authorization_servers")) { + json.at("authorization_servers").get_to(metadata.authorization_servers); + } + if (json.contains("scopes_supported")) { + metadata.scopes_supported = json.at("scopes_supported").get>(); + } +} + +void from_json(const nlohmann::json& json, AuthServerMetadata& metadata) { + metadata.raw = json; + if (json.contains("issuer")) { + json.at("issuer").get_to(metadata.issuer); + } + if (json.contains("authorization_endpoint")) { + json.at("authorization_endpoint").get_to(metadata.authorization_endpoint); + } + if (json.contains("token_endpoint")) { + json.at("token_endpoint").get_to(metadata.token_endpoint); + } + if (json.contains("revocation_endpoint")) { + metadata.revocation_endpoint = json.at("revocation_endpoint").get(); + } + if (json.contains("registration_endpoint")) { + metadata.registration_endpoint = json.at("registration_endpoint").get(); + } + if (json.contains("scopes_supported")) { + metadata.scopes_supported = json.at("scopes_supported").get>(); + } + if (json.contains("response_types_supported")) { + metadata.response_types_supported = + json.at("response_types_supported").get>(); + } + if (json.contains("grant_types_supported")) { + metadata.grant_types_supported = + json.at("grant_types_supported").get>(); + } + if (json.contains("code_challenge_methods_supported")) { + metadata.code_challenge_methods_supported = + json.at("code_challenge_methods_supported").get>(); + } + if (json.contains("authorization_response_iss_parameter_supported")) { + metadata.authorization_response_iss_parameter_supported = + json.at("authorization_response_iss_parameter_supported").get(); + } + if (json.contains("client_id_metadata_document_supported")) { + metadata.client_id_metadata_document_supported = + json.at("client_id_metadata_document_supported").get(); + } + if (json.contains("token_endpoint_auth_methods_supported")) { + metadata.token_endpoint_auth_methods_supported = + json.at("token_endpoint_auth_methods_supported").get>(); + } +} + +struct OAuthDiscoveryClient::Impl { + struct UrlComponents { + std::string scheme; + std::string authority; + std::string path; + }; + + template + struct DiscoveryOperation { + std::shared_ptr owner; + std::string cache_key; + std::vector urls; + std::size_t next_url{0}; + nlohmann::json response; + std::optional metadata; + /// Caller's verdict, consulted before the document is cached or returned. Empty means + /// accept everything that parsed, which is the contract of the overloads that take no + /// acceptor. + std::function accept; + }; + + Impl(std::shared_ptr client, std::chrono::seconds ttl) + : http_client(std::move(client)), cache_ttl(ttl) {} + + static UrlComponents parse_url_components(const std::string& url) { + UrlComponents result; + const auto scheme_end = url.find("://"); + if (scheme_end == std::string::npos) { + // The authorization server identifier a protected-resource document names arrives here + // with nothing having validated it, so this is the throw site a peer reaches first. + throw std::invalid_argument("URL missing scheme: " + sanitize_for_diagnostics(url)); + } + result.scheme = url.substr(0, scheme_end); + auto rest = url.substr(scheme_end + 3); + + const auto path_start = rest.find('/'); + if (path_start == std::string::npos) { + result.authority = std::move(rest); + result.path = "/"; + } else { + result.authority = rest.substr(0, path_start); + result.path = rest.substr(path_start); + } + return result; + } + + template + static Task return_cached(std::shared_ptr metadata) { + co_return std::move(*metadata); + } + + /// Serve a cached document, but only if the caller still accepts it. + /// + /// A coroutine rather than a plain check at the call site so that a rejection surfaces from the + /// await like every other discovery failure, instead of throwing out of the call that merely + /// builds the awaitable. + static Task return_accepted_cached( + std::shared_ptr metadata, ProtectedResourceAcceptor accept) { + if (accept) { + accept(*metadata); + } + co_return std::move(*metadata); + } + + static Task discover_protected_resource( + std::shared_ptr owner, std::string resource_url, + const std::optional& challenge_metadata_url, + ProtectedResourceAcceptor accept = {}) { + // A challenge-supplied URL is its own cache key, so a server that moves its metadata is + // never served an entry discovered through the well-known fallback. + auto cache_key = challenge_metadata_url.value_or(resource_url); + { + std::lock_guard lock(owner->cache_mutex); + const auto iter = owner->resource_cache.find(cache_key); + if (iter != owner->resource_cache.end() && !iter->second.is_expired()) { + return return_accepted_cached( + std::make_shared(iter->second.data), std::move(accept)); + } + } + + auto operation = std::make_shared>(); + operation->owner = std::move(owner); + operation->cache_key = std::move(cache_key); + operation->accept = std::move(accept); + + if (challenge_metadata_url) { + // The challenge named the location; take the server at its word and try nothing else. + operation->urls.push_back(*challenge_metadata_url); + return run_protected_discovery(std::move(operation)); + } + + const auto parsed = parse_url_components(resource_url); + const auto base = parsed.scheme + "://" + parsed.authority; + if (!parsed.path.empty() && parsed.path != "/") { + auto path_part = parsed.path; + if (path_part.front() == '/') { + path_part.erase(0, 1); + } + operation->urls.push_back(base + "/.well-known/oauth-protected-resource/" + path_part); + } + operation->urls.push_back(base + "/.well-known/oauth-protected-resource"); + return run_protected_discovery(std::move(operation)); + } + + static Task discover_auth_server(std::shared_ptr owner, + std::string issuer_url) { + { + std::lock_guard lock(owner->cache_mutex); + const auto iter = owner->auth_cache.find(issuer_url); + if (iter != owner->auth_cache.end() && !iter->second.is_expired()) { + return return_cached(std::make_shared(iter->second.data)); + } + } + + const auto parsed = parse_url_components(issuer_url); + const auto base = parsed.scheme + "://" + parsed.authority; + auto operation = std::make_shared>(); + operation->owner = std::move(owner); + operation->cache_key = std::move(issuer_url); + + if (!parsed.path.empty() && parsed.path != "/") { + auto path_part = parsed.path; + if (path_part.front() == '/') { + path_part.erase(0, 1); + } + if (!path_part.empty() && path_part.back() == '/') { + path_part.pop_back(); + } + operation->urls.push_back(base + "/.well-known/oauth-authorization-server/" + path_part); + operation->urls.push_back(base + "/.well-known/openid-configuration/" + path_part); + operation->urls.push_back(operation->cache_key + "/.well-known/openid-configuration"); + } else { + operation->urls.push_back(base + "/.well-known/oauth-authorization-server"); + operation->urls.push_back(base + "/.well-known/openid-configuration"); + } + return run_auth_discovery(std::move(operation)); + } + + static Task run_protected_discovery( + std::shared_ptr> operation) { + while (operation->next_url < operation->urls.size()) { + try { + const auto index = operation->next_url++; + operation->response = + co_await operation->owner->http_client->get_json(operation->urls[index]); + operation->metadata = operation->response.get(); + } catch (const MetadataPolicyError&) { + // A refused target is a security decision, not a candidate that missed. + throw; + } catch (...) { + continue; + } + + // Outside the try above, because a caller refusing a fetched and parsed document is a + // verdict on it, not a candidate that missed: `catch (...)` would move on to the next + // well-known URL. Before the cache write, because a refused document must never become a + // cache entry served to a later attempt without reaching the network. + if (operation->accept) { + operation->accept(*operation->metadata); + } + + { + std::lock_guard lock(operation->owner->cache_mutex); + operation->owner->resource_cache[operation->cache_key] = { + *operation->metadata, + std::chrono::steady_clock::now() + operation->owner->cache_ttl, + }; + } + co_return *operation->metadata; + } + throw std::runtime_error("Failed to discover protected resource metadata for " + + sanitize_for_diagnostics(operation->cache_key)); + } + + static Task run_auth_discovery( + std::shared_ptr> operation) { + while (operation->next_url < operation->urls.size()) { + try { + const auto index = operation->next_url++; + operation->response = + co_await operation->owner->http_client->get_json(operation->urls[index]); + operation->metadata = operation->response.get(); + + std::lock_guard lock(operation->owner->cache_mutex); + operation->owner->auth_cache[operation->cache_key] = { + *operation->metadata, + std::chrono::steady_clock::now() + operation->owner->cache_ttl, + }; + co_return *operation->metadata; + } catch (const MetadataPolicyError&) { + // A refused target is a security decision, not a candidate that missed. + throw; + } catch (...) { + continue; + } + } + throw std::runtime_error("Failed to discover authorization server metadata for " + + sanitize_for_diagnostics(operation->cache_key)); + } + + std::shared_ptr http_client; + std::chrono::seconds cache_ttl; + std::mutex cache_mutex; + std::unordered_map> resource_cache; + std::unordered_map> auth_cache; +}; + +OAuthDiscoveryClient::OAuthDiscoveryClient(std::shared_ptr http_client, + std::chrono::seconds cache_ttl) + : impl_(std::make_shared(std::move(http_client), cache_ttl)) {} + +Task OAuthDiscoveryClient::discover_protected_resource( + const std::string& resource_url) { + return Impl::discover_protected_resource(impl_, resource_url, std::nullopt); +} + +Task OAuthDiscoveryClient::discover_protected_resource( + const std::string& resource_url, const std::optional& challenge_metadata_url) { + return Impl::discover_protected_resource(impl_, resource_url, challenge_metadata_url); +} + +Task OAuthDiscoveryClient::discover_protected_resource( + const std::string& resource_url, const std::optional& challenge_metadata_url, + ProtectedResourceAcceptor accept) { + return Impl::discover_protected_resource(impl_, resource_url, challenge_metadata_url, + std::move(accept)); +} + +Task OAuthDiscoveryClient::discover_auth_server(const std::string& issuer_url) { + return Impl::discover_auth_server(impl_, issuer_url); +} + +void OAuthDiscoveryClient::clear_cache() { + std::lock_guard lock(impl_->cache_mutex); + impl_->resource_cache.clear(); + impl_->auth_cache.clear(); +} + +namespace { + +struct MiddlewareInvocation { + TokenValidator validator; + mcp::Context* context; + nlohmann::json params; + TypeErasedHandler next; + std::string token; +}; + +Task invoke_auth_middleware(std::shared_ptr invocation) { + // A rejected token is reported to the caller as a tool result rather than raised: throwing from + // middleware is reserved for failures the caller cannot act on, and surfaces as -32603. + if (invocation->token.empty()) { + co_return nlohmann::json(make_tool_error_result("Unauthorized: missing Bearer token")); + } + if (!co_await invocation->validator(invocation->token)) { + co_return nlohmann::json(make_tool_error_result("Unauthorized: invalid Bearer token")); + } + co_return co_await invocation->next(*invocation->context, invocation->params); +} + +} // namespace + +Middleware make_auth_middleware(TokenValidator validator) { + return [validator = std::move(validator)](mcp::Context& context, const nlohmann::json& params, + TypeErasedHandler next) -> Task { + auto invocation = std::make_shared(); + invocation->validator = validator; + invocation->context = &context; + invocation->params = params; + invocation->next = std::move(next); + if (params.contains("_meta") && params["_meta"].contains("auth_token")) { + invocation->token = params["_meta"]["auth_token"].get(); + } + return invoke_auth_middleware(std::move(invocation)); + }; +} + +std::string extract_bearer_token(std::string_view auth_header_value) { + constexpr std::string_view prefix = "Bearer "; + if (auth_header_value.size() > prefix.size() && + auth_header_value.substr(0, prefix.size()) == prefix) { + return std::string(auth_header_value.substr(prefix.size())); + } + return {}; +} + +struct OAuthAuthenticator::Impl { + struct RefreshOperation { + std::shared_ptr owner; + TokenResponse stored_token; + std::optional new_token; + }; + + Impl(std::shared_ptr store, std::shared_ptr client, + OAuthConfig oauth_config, std::string url) + : token_store(std::move(store)), + oauth_client(std::move(client)), + scope(oauth_client->make_scope()), + config(std::move(oauth_config)), + server_url(std::move(url)) {} + + static Task return_false() { co_return false; } + + static Task try_refresh(std::shared_ptr owner) { + auto stored = owner->token_store->load(owner->server_url); + if (!stored || !stored->refresh_token) { + return return_false(); + } + + auto operation = std::make_shared(); + operation->owner = std::move(owner); + operation->stored_token = std::move(*stored); + return run_refresh(std::move(operation)); + } + + static Task run_refresh(std::shared_ptr operation) { + try { + operation->new_token = co_await operation->owner->scope.refresh_token( + operation->owner->config, *operation->stored_token.refresh_token); + if (!operation->new_token->refresh_token) { + operation->new_token->refresh_token = operation->stored_token.refresh_token; + } + operation->owner->token_store->store(operation->owner->server_url, + std::move(*operation->new_token)); + co_return true; + } catch (...) { + co_return false; + } + } + + std::shared_ptr token_store; + /// Kept so the client outlives this authenticator, even though every request goes through the + /// scope below. + std::shared_ptr oauth_client; + /// This authenticator's own slice of the client. The client is supplied by the application and + /// may be shared with other authenticators for other servers, so close() must end this + /// authenticator's work without ending theirs. + OAuthHttpClientScope scope; + OAuthConfig config; + std::string server_url; +}; + +namespace { + +/// Impl's constructor takes a scope off the client, so a null one faults before the object exists. +/// Refuse it the way OAuthClientTransport's constructor refuses its own null arguments. +std::shared_ptr require_http_client(std::shared_ptr client) { + if (!client) { + throw std::invalid_argument("OAuthAuthenticator requires an OAuth HTTP client"); + } + return client; +} + +} // namespace + +OAuthAuthenticator::OAuthAuthenticator(std::shared_ptr token_store, + std::shared_ptr oauth_client, + OAuthConfig config, std::string server_url) + : impl_(std::make_shared(std::move(token_store), require_http_client(std::move(oauth_client)), + std::move(config), std::move(server_url))) {} + +std::string OAuthAuthenticator::get_access_token() const { + const auto token = impl_->token_store->load(impl_->server_url); + return token ? token->access_token : std::string{}; +} + +Task OAuthAuthenticator::try_refresh_token() { return Impl::try_refresh(impl_); } + +void OAuthAuthenticator::store_token(TokenResponse token) { + impl_->token_store->store(impl_->server_url, std::move(token)); +} + +/// Ends only this authenticator's requests. Do not reach for abort_pending() here: it latches the +/// whole client irreversibly, which is correct for a client its owner created and wrong for one the +/// application supplied and may share. Two authenticators for two servers on one client would then +/// silently disable each other -- run_refresh() reports a failed refresh as a plain `false`, so the +/// survivor gets no error, only tokens that quietly stop renewing. +void OAuthAuthenticator::close() { impl_->scope.abort(); } + +namespace { + +/// Join scopes into the space-delimited form an authorization request carries. +std::string join_scopes(const std::vector& scopes) { + std::string joined; + for (const auto& scope : scopes) { + if (!joined.empty()) { + joined.push_back(' '); + } + joined += scope; + } + return joined; +} + +std::vector split_scopes(std::string_view scopes) { + std::vector split; + std::size_t cursor = 0; + while (cursor < scopes.size()) { + const auto next = scopes.find(' ', cursor); + const auto token = scopes.substr(cursor, next == std::string_view::npos ? next : next - cursor); + if (!token.empty()) { + split.emplace_back(token); + } + if (next == std::string_view::npos) { + break; + } + cursor = next + 1; + } + return split; +} + +/// Union of two scope strings, preserving the order of `primary` and appending whatever only +/// `secondary` carries. Step-up authorization re-authorizes on a challenge that names only what the +/// refused operation needed, so without the union an earlier grant's scopes would be dropped. +std::optional union_scopes(const std::optional& primary, + const std::optional& secondary) { + if (!primary || primary->empty()) { + return secondary && !secondary->empty() ? secondary : std::nullopt; + } + if (!secondary || secondary->empty()) { + return primary; + } + auto merged = split_scopes(*primary); + for (auto& candidate : split_scopes(*secondary)) { + if (std::find(merged.begin(), merged.end(), candidate) == merged.end()) { + merged.push_back(std::move(candidate)); + } + } + return join_scopes(merged); +} + +/// Choose the token endpoint authentication method from what the server advertises. +/// +/// A server that publishes the list has told the client which methods it will accept, so the +/// client picks the strongest one it can actually satisfy rather than guessing. A server that +/// publishes nothing leaves the decision unset, which sends the secret, when there is one, in the +/// request body. +std::optional select_token_endpoint_auth_method( + const std::optional>& supported, bool has_client_secret) { + if (!supported || supported->empty()) { + return std::nullopt; + } + const auto advertises = [&supported](std::string_view method) { + return std::find(supported->begin(), supported->end(), method) != supported->end(); + }; + if (has_client_secret) { + if (advertises("client_secret_basic")) { + return "client_secret_basic"; + } + if (advertises("client_secret_post")) { + return "client_secret_post"; + } + } + if (advertises("none")) { + return "none"; + } + return std::nullopt; +} + +std::string append_query(const std::string& endpoint, const std::string& query) { + if (query.empty()) { + return endpoint; + } + const auto separator = endpoint.find('?') == std::string::npos ? '?' : '&'; + return endpoint + separator + query; +} + +} // namespace + +struct OAuthAuthorizationManager::Impl { + /// One coalescing authorization attempt, and the channel every follower of it reads its result + /// from. + /// + /// The outcome lives on the attempt, not the manager, so a follower that wakes after a later + /// flight finished cannot read that flight's result as its own. It carries the exception rather + /// than a bool so a follower sees the leader's diagnostic. `succeeded` and `failure` are written + /// once by the leader in run_leading_challenge() and read by followers after they wake, both + /// under `state_mutex`; `expire_flight()` wakes them and is always called after that write. + struct Flight { + explicit Flight(const net::strand& flight_strand) + : timer(flight_strand, net::steady_timer::time_point::max()) {} + + /// Never waited on to expire naturally: it parks at time_point::max() and is pushed into + /// the past to release the followers. See expire_flight(). + net::steady_timer timer; + /// Set when the leader has recorded its result below. A follower woken by close() rather + /// than by its leader finishing sees this false. + bool finished{false}; + bool succeeded{false}; + /// The leader's own exception, rethrown by every follower so they all report the same + /// reason the leader does. Null when the flight finished without throwing. + std::exception_ptr failure; + }; + + struct ChallengeOperation { + std::shared_ptr owner; + BearerChallenge challenge; + std::optional resource_metadata; + std::optional auth_metadata; + ClientIdentityServerFacts facts; + std::optional stored_identity; + OAuthClientInformation identity; + nlohmann::json registration_response; + AuthorizationRequest request; + AuthorizationResponse response; + OAuthConfig token_config; + TokenResponse token; + }; + + struct RefreshOperation { + std::shared_ptr owner; + OAuthConfig token_config; + TokenResponse stored_token; + std::optional new_token; + }; + + Impl(const net::any_io_executor& executor, std::shared_ptr store, + OAuthAuthorizationConfig authorization_config, AuthorizationCallback authorization_callback) + : token_store(std::move(store)), + http_client(std::make_shared(executor)), + flight_strand(net::make_strand(executor)), + config(std::move(authorization_config)), + callback(std::move(authorization_callback)) { + // A bare `client_id` is the shorthand form of injected credentials, so the two spellings + // reach the same terminal decision instead of one of them quietly permitting registration. + if (!config.client_identity.pre_registered && !config.client_id.empty()) { + OAuthClientInformation injected; + injected.client_id = config.client_id; + injected.client_secret = config.client_secret; + // Without this the issuer binding in `select_client_identity` is inert on the + // shorthand path: an empty issuer would make every authorization server look like + // the one these credentials belong to. + injected.issuer = config.client_issuer; + injected.source = ClientIdentitySource::pre_registered; + config.client_identity.pre_registered = std::move(injected); + } + if (config.client_identity.metadata.redirect_uris.empty() && !config.redirect_uri.empty()) { + config.client_identity.metadata.redirect_uris.push_back(config.redirect_uri); + } + // One call rather than two, even here where the client has not yet issued a request: the + // pair method is what a reader should find at a site that installs both. An empty + // `host_resolver` leaves this client on the executor's system resolver. + http_client->configure(config.policy, config.host_resolver); + discovery = std::make_shared(http_client); + } + + static Task return_false() { co_return false; } + + /// Fails the same way return_false() succeeds: lazily, when the returned awaitable is awaited. + /// handle_challenge() is a plain function returning an awaitable, so a bare `throw` in its body + /// would fire when try_handle_challenge() is *called*, unlike the virtual it overrides, whose + /// contract is the lazy one. It stays a plain function because making it a coroutine would add a + /// frame and a suspension to a path that runs on every 401. + static Task throw_closed() { + throw std::runtime_error("OAuth authorization manager closed"); + co_return false; + } + + /// Challenge scope is authoritative; `scopes_supported` is the fallback; otherwise no scope is + /// requested at all. An explicit configured scope overrides both. The two sources are never + /// merged with each other: the spec forbids assuming any set relationship between them. Scope + /// already granted by an earlier authorization *is* merged in, because a 403 step-up challenge + /// names only what the refused operation needed. + static std::optional select_scope(const Impl& owner, const BearerChallenge& challenge, + const ProtectedResourceMetadata& resource) { + if (owner.config.scope) { + return owner.config.scope; + } + std::optional selected; + if (challenge.scope && !challenge.scope->empty()) { + selected = challenge.scope; + } else if (resource.scopes_supported && !resource.scopes_supported->empty()) { + selected = join_scopes(*resource.scopes_supported); + } + std::optional granted; + { + std::lock_guard lock(owner.state_mutex); + granted = owner.granted_scope; + } + return union_scopes(selected, granted); + } + + /// Record the four authorization-server facts client identity selection turns on. + static ClientIdentityServerFacts server_facts(const AuthServerMetadata& auth_server) { + ClientIdentityServerFacts facts; + facts.issuer = auth_server.issuer; + facts.client_id_metadata_document_supported = + auth_server.client_id_metadata_document_supported.value_or(false); + facts.registration_endpoint = auth_server.registration_endpoint; + if (auth_server.scopes_supported) { + facts.scopes_supported = *auth_server.scopes_supported; + } + return facts; + } + + /// Load the credentials held for this issuer, discarding any entry that is not usable for it. + static std::optional load_stored_identity( + const Impl& owner, const ClientIdentityServerFacts& facts) { + if (!owner.config.credential_store || facts.issuer.empty()) { + return std::nullopt; + } + auto stored = owner.config.credential_store->load(facts.issuer); + if (!stored) { + return std::nullopt; + } + // A stored entry that disagrees with its own key, or whose secret the server has already + // retired, is worse than no entry at all: presenting it would either misbind the + // credential or fail the exchange with a stale one. + const auto now = static_cast(std::time(nullptr)); + if (stored->issuer != facts.issuer || stored->secret_expired(now)) { + return std::nullopt; + } + return stored; + } + + static Task resolve_identity( + std::shared_ptr operation) { + auto& owner = *operation->owner; + const auto decision = select_client_identity(owner.config.client_identity, operation->facts, + operation->stored_identity); + switch (decision) { + case ClientIdentityDecision::use_pre_registered: { + auto identity = *owner.config.client_identity.pre_registered; + identity.source = ClientIdentitySource::pre_registered; + identity.issuer = operation->facts.issuer; + co_return identity; + } + case ClientIdentityDecision::use_client_id_metadata_document: { + // The document URL is the client identifier itself, so there is nothing to + // register and nothing to persist. + OAuthClientInformation identity; + identity.client_id = *owner.config.client_identity.client_metadata_url; + identity.issuer = operation->facts.issuer; + identity.source = ClientIdentitySource::client_id_metadata_document; + co_return identity; + } + case ClientIdentityDecision::reuse_stored_registration: + co_return *operation->stored_identity; + case ClientIdentityDecision::register_dynamically: + break; + case ClientIdentityDecision::unavailable: { + // `unavailable` covers several causes; say which one, because the unbound-secret + // refusal is a configuration mistake the caller can fix and the generic message + // would send them looking in the wrong place. + const auto& injected = owner.config.client_identity.pre_registered; + if (injected && !injected->client_id.empty() && injected->issuer.empty() && + injected->client_secret && !injected->client_secret->empty()) { + // Both spellings are named, with the condition on each, because they are not + // interchangeable: the constructor copies `client_issuer` into the injected + // credentials only when `client_identity.pre_registered` was not already set, so + // a caller who built that struct themselves can set `client_issuer` and watch it + // be ignored. + throw std::runtime_error( + "Injected client credentials carry a client_secret but name no issuer, so " + "they cannot be presented to authorization server " + + sanitize_for_diagnostics(operation->facts.issuer) + + "; set the issuer these credentials are bound to. Set " + "client_identity.pre_registered.issuer if you populated " + "client_identity.pre_registered yourself; client_issuer applies only to " + "credentials given as client_id and client_secret, and is ignored once " + "client_identity.pre_registered is set"); + } + throw std::runtime_error("No client identity is available for authorization server " + + sanitize_for_diagnostics(operation->facts.issuer)); + } + } + + operation->registration_response = co_await owner.http_client->post_json( + *operation->facts.registration_endpoint, + build_registration_request(owner.config.client_identity.metadata, operation->facts)); + + auto registered = operation->registration_response.get(); + if (registered.client_id.empty()) { + throw std::runtime_error("Client registration response omitted client_id"); + } + registered.issuer = operation->facts.issuer; + registered.source = ClientIdentitySource::dynamic_registration; + if (owner.config.credential_store) { + owner.config.credential_store->store(operation->facts.issuer, registered); + } + co_return registered; + } + + static AuthorizationRequest build_request(const Impl& owner, const BearerChallenge& challenge, + const ProtectedResourceMetadata& resource, + const AuthServerMetadata& auth_server, + const OAuthClientInformation& identity) { + const auto pkce = generate_pkce_pair(); + + AuthorizationRequest request; + request.state = mcp::detail::generate_secure_session_id(); + request.code_verifier = pkce.code_verifier; + request.code_challenge = pkce.code_challenge; + // Recorded from the metadata document this client fetched itself. Response validation is + // only as trustworthy as this value's provenance. + request.issuer = auth_server.issuer; + request.issuer_parameter_supported = + auth_server.authorization_response_iss_parameter_supported.value_or(false); + request.client_id = identity.client_id; + request.redirect_uri = owner.config.redirect_uri; + request.scope = select_scope(owner, challenge, resource); + // Defence in depth: `run_challenge` already rejected an unidentified resource before any + // outbound request. Re-checking here is a pure string comparison that cannot fail on that + // path, and keeps `build_request` correct if a second caller ever appears. + require_resource_identifies_server(resource.resource, owner.config.server_url); + request.resource = resource.resource; + + KeyValuePairList params = { + {"response_type", "code"}, + {"client_id", request.client_id}, + {"redirect_uri", request.redirect_uri}, + {"state", request.state}, + {"code_challenge", request.code_challenge}, + {"code_challenge_method", pkce.challenge_method}, + }; + if (request.scope) { + params.emplace_back("scope", *request.scope); + } + // RFC 8707: the resource indicator travels on the authorization request as well as the + // token request. + if (request.resource) { + params.emplace_back("resource", *request.resource); + } + request.authorization_url = + append_query(auth_server.authorization_endpoint, detail::build_form_body(params)); + return request; + } + + static OAuthConfig build_token_config(const Impl& owner, const AuthServerMetadata& auth_server, + const AuthorizationRequest& request, + const OAuthClientInformation& identity) { + OAuthConfig token_config; + token_config.client_id = identity.client_id; + token_config.client_secret = identity.client_secret; + token_config.token_endpoint = auth_server.token_endpoint; + token_config.authorization_endpoint = auth_server.authorization_endpoint; + token_config.revocation_endpoint = auth_server.revocation_endpoint; + token_config.redirect_uri = owner.config.redirect_uri; + token_config.scope = request.scope; + token_config.resource = request.resource; + token_config.token_endpoint_auth_method = select_token_endpoint_auth_method( + auth_server.token_endpoint_auth_methods_supported, identity.client_secret.has_value()); + return token_config; + } + + static Task run_challenge(std::shared_ptr operation) { + auto& owner = *operation->owner; + + // Both checks that decide whether this document may be acted on, handed to discovery as its + // acceptor so they gate the cache write and also run on a cache hit: applied afterwards, a + // refused document would already be cached for the full TTL. Order matters: a document that + // both lists no authorization servers and carries a resource that is not ours must report the + // missing authorization servers. + auto accept = [owner = operation->owner](const ProtectedResourceMetadata& resource) { + if (resource.authorization_servers.empty()) { + throw std::runtime_error("Protected resource metadata listed no authorization servers"); + } + // Refused before the first outbound request this document would drive, so a PRM that + // does not identify our configured server cannot make us fetch the authorization server + // metadata it names, register a client with that server, or persist those credentials. + require_resource_identifies_server(resource.resource, owner->config.server_url); + }; + + // The challenge's own metadata URL wins over the well-known fallback order. + operation->resource_metadata = co_await owner.discovery->discover_protected_resource( + owner.config.server_url, operation->challenge.resource_metadata, std::move(accept)); + + operation->auth_metadata = co_await owner.discovery->discover_auth_server( + operation->resource_metadata->authorization_servers.front()); + // RFC 8414 §3.3: the issuer a metadata document claims MUST be the location it was fetched + // from, compared byte-exact. Normalizing here would re-open the AS mix-up this closes. + if (operation->auth_metadata->issuer.empty() || + operation->auth_metadata->issuer != + operation->resource_metadata->authorization_servers.front()) { + throw std::runtime_error( + "Authorization server metadata issuer does not identify the server it was fetched " + "from"); + } + if (operation->auth_metadata->authorization_endpoint.empty() || + operation->auth_metadata->token_endpoint.empty()) { + throw std::runtime_error( + "Authorization server metadata omitted an authorization or token endpoint"); + } + + operation->facts = server_facts(*operation->auth_metadata); + operation->stored_identity = load_stored_identity(owner, operation->facts); + operation->identity = co_await resolve_identity(operation); + + operation->request = build_request(owner, operation->challenge, *operation->resource_metadata, + *operation->auth_metadata, operation->identity); + operation->token_config = build_token_config(owner, *operation->auth_metadata, + operation->request, operation->identity); + { + std::lock_guard lock(owner.state_mutex); + owner.last_request = operation->request; + owner.last_identity = operation->identity; + owner.token_config = operation->token_config; + } + + operation->response = co_await owner.callback(operation->request); + const auto validation = + validate_authorization_response(operation->request, operation->response); + if (!validation.accepted()) { + throw std::runtime_error("Authorization response rejected: " + validation.message); + } + + operation->token = co_await owner.http_client->exchange_code( + operation->token_config, *operation->response.code, operation->request.code_verifier); + { + std::lock_guard lock(owner.state_mutex); + // What the server actually granted, falling back to what was asked for when the token + // response stays silent (RFC 6749 §5.1 makes `scope` optional in that case). + owner.granted_scope = + operation->token.scope ? operation->token.scope : operation->request.scope; + } + owner.token_store->store(owner.config.server_url, operation->token); + co_return true; + } + + /// Push `flight`'s deadline into the past, on `flight_strand` so it never races the timer's own + /// operations: expires_at() and async_wait() both execute on the calling thread (the timer's + /// executor governs only its completion handler), and await_in_flight() initiates its wait on the + /// same strand. + /// + /// expires_at() rather than cancel(): it also moves the deadline, so a follower that has joined + /// `flight` but not yet called async_wait() completes immediately instead of parking on + /// time_point::max() forever. A null `flight` is a no-op. + static void expire_flight(const net::strand& flight_strand, + std::shared_ptr flight) { + if (!flight) { + return; + } + net::dispatch(flight_strand, [flight = std::move(flight)]() { + flight->timer.expires_at(net::steady_timer::time_point::min()); + }); + } + + /// Waiters share the leader's outcome instead of opening a second authorization flow, so a + /// burst of concurrent requests that all hit the same challenge authorizes exactly once. + static Task await_in_flight(std::shared_ptr owner, std::shared_ptr flight) { + // async_wait() touches the timer synchronously on the thread that calls it, so it has to be + // initiated on the same strand expire_flight() dispatches its expires_at() onto; the timer's + // associated executor governs only where its completion handler runs. + // + // bind_executor() rather than a bare post(flight_strand, use_awaitable): a bare post leaves + // the resumption on this coroutine's own executor and only happens to land inside the strand + // when the two share one io_context. + co_await net::dispatch(net::bind_executor(owner->flight_strand, net::use_awaitable)); + boost::system::error_code ignored; + co_await flight->timer.async_wait(net::redirect_error(net::use_awaitable, ignored)); + + // The dispatch above moves this coroutine onto flight_strand, but only until the next + // suspension: an awaitable's executor is fixed when it is spawned, and the wait's completion + // handler carries that executor, so the caller is back on its own executor here and no + // explicit hop back is needed. That matters -- Client spawns its write onto its own strand + // and SerializedTransportWriter builds another to serialise writes, and both would be + // bypassed by a continuation left on flight_strand. Asserted by + // AFollowerReleasedFromTheSingleFlightTimerResumesOnItsOwnStrand. + std::exception_ptr failure; + bool succeeded = false; + { + std::lock_guard lock(owner->state_mutex); + // close() expires the same timer to release a follower parked here; a follower that + // wakes because the manager closed gets a clear error rather than the misleading "not + // authorized" that the leader's own outcome would otherwise report. Checked before the + // flight's own result because a closed manager is the more specific answer: the flight + // it joined may never have finished at all. + if (owner->closed) { + throw std::runtime_error("OAuth authorization manager closed"); + } + // Read off the flight this follower actually waited on. Reading a manager-wide field + // here would let a late waker report a different attempt's outcome. + // + // `finished` is a guard rather than a case that arises today: the only two things that + // expire this timer are the leader recording its result and close(), and the closed check + // above already took the second. It keeps a wake-up added later from making a follower + // report a flat "not authorized" that is indistinguishable from a real refusal. + if (!flight->finished) { + throw std::runtime_error( + "OAuth authorization attempt ended without recording an outcome"); + } + failure = flight->failure; + succeeded = flight->succeeded; + } + + // Rethrown outside the lock: the leader's exception must not travel through a destructor + // while state_mutex is held. Every follower rethrows the same exception_ptr, which is safe -- + // the object it refers to is shared and read-only. Nothing on the follower path swallows it: + // run_write() catches (...) only to erase_pending() and rethrow. + if (failure) { + std::rethrow_exception(failure); + } + co_return succeeded; + } + + static Task run_leading_challenge(std::shared_ptr operation) { + auto owner = operation->owner; + bool succeeded = false; + std::exception_ptr failure; + try { + succeeded = co_await run_challenge(std::move(operation)); + } catch (...) { + failure = std::current_exception(); + } + + std::shared_ptr finished; + { + std::lock_guard lock(owner->state_mutex); + finished = std::move(owner->flight); + owner->flight.reset(); + if (finished) { + // Recorded on the attempt itself, before any follower is woken, so each follower + // reads the result of the flight it joined and gets the leader's own reason for it + // rather than a bare false. The same exception_ptr is rethrown below. + finished->succeeded = succeeded; + finished->failure = failure; + finished->finished = true; + } + } + // A late joiner may have read the old `flight` out of the lock above just before this reset + // and not yet be waiting on it (see expire_flight()'s comment); expires_at(), not cancel(), + // is what still reaches it. + expire_flight(owner->flight_strand, std::move(finished)); + if (failure) { + std::rethrow_exception(failure); + } + co_return succeeded; + } + + static Task handle_challenge(std::shared_ptr owner, const std::string& header) { + auto challenge = select_bearer_challenge(parse_www_authenticate(header)); + if (!challenge) { + return return_false(); + } + + std::shared_ptr joined; + { + std::lock_guard lock(owner->state_mutex); + // Checked in the same critical section that reads or creates `flight`: close() also + // takes this lock, so the two can never interleave as a new flight being created right + // after close() already ran past it, unnoticed. (http_client's own sticky abort would + // still stop that flight's first network call either way, but this keeps a closed + // manager from starting one at all.) + if (owner->closed) { + return throw_closed(); + } + if (owner->flight) { + joined = owner->flight; + } else { + owner->flight = std::make_shared(owner->flight_strand); + } + } + if (joined) { + return await_in_flight(std::move(owner), std::move(joined)); + } + + auto operation = std::make_shared(); + operation->owner = std::move(owner); + operation->challenge = std::move(*challenge); + return run_leading_challenge(std::move(operation)); + } + + static Task try_refresh(std::shared_ptr owner) { + auto stored = owner->token_store->load(owner->config.server_url); + if (!stored || !stored->refresh_token) { + return return_false(); + } + + auto operation = std::make_shared(); + { + std::lock_guard lock(owner->state_mutex); + if (!owner->token_config) { + return return_false(); + } + operation->token_config = *owner->token_config; + } + operation->owner = std::move(owner); + operation->stored_token = std::move(*stored); + return run_refresh(std::move(operation)); + } + + static Task run_refresh(std::shared_ptr operation) { + try { + operation->new_token = co_await operation->owner->http_client->refresh_token( + operation->token_config, *operation->stored_token.refresh_token); + if (!operation->new_token->refresh_token) { + operation->new_token->refresh_token = operation->stored_token.refresh_token; + } + operation->owner->token_store->store(operation->owner->config.server_url, + std::move(*operation->new_token)); + co_return true; + } catch (const MetadataPolicyError&) { + throw; + } catch (...) { + co_return false; + } + } + + std::shared_ptr token_store; + std::shared_ptr http_client; + std::shared_ptr discovery; + /// Serialises every access to `flight`. A boost::asio::steady_timer is not safe for concurrent + /// use, and its two touch points -- expire_flight()'s expires_at() and await_in_flight()'s + /// async_wait() -- both execute synchronously on whatever thread calls them, so nothing but a + /// strand shared by both of them keeps them apart on a multi-threaded io_context. The timer is + /// constructed on this strand as well, so its completion handlers run here too. + net::strand flight_strand; + OAuthAuthorizationConfig config; + AuthorizationCallback callback; + mutable std::mutex state_mutex; + std::optional last_request; + std::optional last_identity; + std::optional token_config; + std::optional granted_scope; + /// Non-null exactly while one authorization flow is running; cancelling it releases the + /// requests that coalesced onto it. + std::shared_ptr flight; + bool closed{false}; + + /// Abort whatever this manager has in flight and release every parked follower with an error. + /// + /// `http_client->abort_pending()` unblocks a leader parked in an HTTP call (discovery, + /// registration and token exchange share the one client) and, being sticky, fails a leader or a + /// racing fresh flow at its first network call. `expire_flight()` releases followers directly, + /// which also covers a leader parked in the application's consent callback, where there is no + /// network state to cancel. + static void close(const std::shared_ptr& owner) { + std::shared_ptr flight; + { + std::lock_guard lock(owner->state_mutex); + if (owner->closed) { + return; + } + owner->closed = true; + flight = owner->flight; + } + owner->http_client->abort_pending(); + expire_flight(owner->flight_strand, std::move(flight)); + } +}; + +OAuthAuthorizationManager::OAuthAuthorizationManager(const net::any_io_executor& executor, + std::shared_ptr token_store, + OAuthAuthorizationConfig config, + AuthorizationCallback callback) { + if (!token_store || !callback) { + throw std::invalid_argument( + "OAuthAuthorizationManager requires a token store and an authorization callback"); + } + impl_ = std::make_shared(executor, std::move(token_store), std::move(config), + std::move(callback)); +} + +std::string OAuthAuthorizationManager::get_access_token() const { + const auto token = impl_->token_store->load(impl_->config.server_url); + return token ? token->access_token : std::string{}; +} + +Task OAuthAuthorizationManager::try_refresh_token() { return Impl::try_refresh(impl_); } + +Task OAuthAuthorizationManager::try_handle_challenge(const std::string& www_authenticate) { + return Impl::handle_challenge(impl_, www_authenticate); +} + +void OAuthAuthorizationManager::close() { Impl::close(impl_); } + +std::optional OAuthAuthorizationManager::last_authorization_request() const { + std::lock_guard lock(impl_->state_mutex); + return impl_->last_request; +} + +std::optional OAuthAuthorizationManager::last_client_identity() const { + std::lock_guard lock(impl_->state_mutex); + return impl_->last_identity; +} + +struct OAuthClientTransport::Impl { + using Clock = std::chrono::steady_clock; + using PendingOrder = std::list; + + struct PendingRequest { + std::shared_ptr wire; + Clock::time_point expires_at; + PendingOrder::iterator order_position; + std::size_t retained_bytes{0}; + std::uint64_t generation{0}; + }; + + using PendingMap = std::unordered_map; + + struct WriteOperation { + std::shared_ptr owner; + std::shared_ptr original; + std::shared_ptr outgoing; + std::optional request_key; + std::optional request_generation; + std::exception_ptr write_error; + std::string authenticate_challenge; + bool authentication_challenge{false}; + bool refresh_allowed{false}; + int authorization_attempts{0}; + }; + + /// At most three authorization challenges are honoured per logical request, so a server that + /// keeps refusing a scope it will never grant cannot drive an unbounded authorization loop. + static constexpr int g_max_authorization_challenges = 3; + + explicit Impl(std::shared_ptr wrapped, + std::shared_ptr token_authenticator, + OAuthClientTransportOptions replay_options) + : inner(std::move(wrapped)), + authenticator(std::move(token_authenticator)), + options(std::move(replay_options)) {} + + static void erase_pending_locked(Impl& owner, PendingMap::iterator iter) { + owner.pending_bytes -= iter->second.retained_bytes; + owner.pending_order.erase(iter->second.order_position); + owner.pending_requests.erase(iter); + } + + static void prune_expired_locked(Impl& owner, Clock::time_point now) { + while (!owner.pending_order.empty()) { + const auto iter = owner.pending_requests.find(owner.pending_order.front()); + if (iter == owner.pending_requests.end()) { + owner.pending_order.pop_front(); + continue; + } + if (iter->second.expires_at > now) { + break; + } + erase_pending_locked(owner, iter); + } + } + + static void evict_oldest_locked(Impl& owner) { + if (owner.pending_order.empty()) { + return; + } + const auto iter = owner.pending_requests.find(owner.pending_order.front()); + if (iter == owner.pending_requests.end()) { + owner.pending_order.pop_front(); + return; + } + erase_pending_locked(owner, iter); + } + + static std::optional retained_size(const Impl& owner, std::string_view key, + std::string_view wire) { + auto remaining = owner.options.max_pending_request_bytes; + if (wire.size() > remaining) { + return std::nullopt; + } + remaining -= wire.size(); + for (int key_copy = 0; key_copy < 2; ++key_copy) { + if (key.size() > remaining) { + return std::nullopt; + } + remaining -= key.size(); + } + return owner.options.max_pending_request_bytes - remaining; + } + + static std::optional remember_request( + const std::shared_ptr& owner, const std::string& key, + const std::shared_ptr& wire) { + std::lock_guard lock(owner->pending_mutex); + const auto now = Clock::now(); + prune_expired_locked(*owner, now); + + if (owner->closed || owner->options.max_pending_requests == 0 || + owner->options.pending_request_ttl <= std::chrono::milliseconds::zero()) { + return std::nullopt; + } + const auto bytes = retained_size(*owner, key, *wire); + if (!bytes) { + return std::nullopt; + } + + if (const auto existing = owner->pending_requests.find(key); + existing != owner->pending_requests.end()) { + erase_pending_locked(*owner, existing); + } + while (owner->pending_requests.size() >= owner->options.max_pending_requests || + *bytes > owner->options.max_pending_request_bytes - owner->pending_bytes) { + evict_oldest_locked(*owner); + } + + ++owner->next_generation; + if (owner->next_generation == 0) { + ++owner->next_generation; + } + const auto generation = owner->next_generation; + const auto max_ttl = + std::chrono::duration_cast(Clock::time_point::max() - now); + const auto expires_at = owner->options.pending_request_ttl >= max_ttl + ? Clock::time_point::max() + : now + owner->options.pending_request_ttl; + + owner->pending_order.push_back(key); + try { + owner->pending_requests.emplace( + key, PendingRequest{wire, expires_at, std::prev(owner->pending_order.end()), *bytes, + generation}); + } catch (...) { + owner->pending_order.pop_back(); + throw; + } + owner->pending_bytes += *bytes; + return generation; + } + + static std::string inject_token(const std::shared_ptr& owner, JSONRPCRequest request) { + const auto token = owner->authenticator->get_access_token(); + if (token.empty()) { + return nlohmann::json(request).dump(); + } + if (!request.params) { + request.params = nlohmann::json::object(); + } + auto& params = *request.params; + if (!params.is_object()) { + params = nlohmann::json::object(); + } + if (!params.contains("_meta")) { + params["_meta"] = nlohmann::json::object(); + } + params["_meta"]["auth_token"] = token; + return nlohmann::json(request).dump(); + } + + static std::shared_ptr prepare_write(std::shared_ptr owner, + std::shared_ptr original) { + auto operation = std::make_shared(); + operation->owner = std::move(owner); + operation->original = std::move(original); + auto outgoing = std::make_shared(*operation->original); + + try { + const auto json_message = nlohmann::json::parse(*operation->original); + const auto message = json_message.get(); + if (const auto* request = std::get_if(&message)) { + operation->request_key = request->id.correlation_key(); + if (!operation->owner->uses_http_authorization_header) { + *outgoing = inject_token(operation->owner, *request); + } + } + } catch (const std::exception&) { + // Non-JSON messages pass through unchanged. + } + + operation->outgoing = std::move(outgoing); + if (operation->request_key) { + operation->request_generation = + remember_request(operation->owner, *operation->request_key, operation->original); + } + return operation; + } + + static Task write(std::shared_ptr owner, std::shared_ptr original) { + return run_write(prepare_write(std::move(owner), std::move(original))); + } + + static void erase_pending(const std::shared_ptr& operation) { + if (!operation->request_key || !operation->request_generation) { + return; + } + std::lock_guard lock(operation->owner->pending_mutex); + const auto iter = operation->owner->pending_requests.find(*operation->request_key); + if (iter != operation->owner->pending_requests.end() && + iter->second.generation == *operation->request_generation) { + erase_pending_locked(*operation->owner, iter); + } + } + + static Task run_write(std::shared_ptr operation) { + while (true) { + operation->write_error = nullptr; + operation->authentication_challenge = false; + operation->refresh_allowed = false; + operation->authenticate_challenge.clear(); + + try { + co_await operation->owner->inner->write_message(*operation->outgoing); + } catch (const mcp::HttpStatusError& error) { + // 401 says the request was not authenticated at all. 403 that carries a challenge + // is the step-up case: the token is valid but its scope does not cover this + // operation, so the challenge drives a fresh authorization rather than a refresh. + const auto unauthorized = + error.status() == static_cast(http::status::unauthorized); + const auto stepped_up = + error.status() == static_cast(http::status::forbidden) && + !error.authenticate_challenge().empty(); + operation->authentication_challenge = + operation->owner->uses_http_authorization_header && (unauthorized || stepped_up); + if (operation->authentication_challenge) { + operation->authenticate_challenge = error.authenticate_challenge(); + operation->refresh_allowed = unauthorized; + } + operation->write_error = std::current_exception(); + } catch (...) { + operation->write_error = std::current_exception(); + } + + if (!operation->write_error) { + co_return; + } + if (!operation->authentication_challenge || + operation->authorization_attempts >= g_max_authorization_challenges) { + break; + } + ++operation->authorization_attempts; + + bool authorized = false; + try { + // A challenge that names its metadata can drive a full authorization exchange; + // renewing an existing grant is the fallback when it cannot. Renewal is pointless + // against an insufficient-scope refusal, so it is offered only for a 401. + if (!operation->authenticate_challenge.empty()) { + authorized = co_await operation->owner->authenticator->try_handle_challenge( + operation->authenticate_challenge); + } + if (!authorized && operation->refresh_allowed) { + authorized = co_await operation->owner->authenticator->try_refresh_token(); + } + } catch (...) { + erase_pending(operation); + throw; + } + if (!authorized) { + break; + } + } + + erase_pending(operation); + std::rethrow_exception(operation->write_error); + } + + static std::shared_ptr take_retry_wire(const std::shared_ptr& owner, + std::string_view raw) { + try { + const auto json_message = nlohmann::json::parse(raw); + const auto message = json_message.get(); + + std::optional key; + bool unauthorized = false; + if (const auto* result = std::get_if(&message)) { + key = result->id.correlation_key(); + } else if (const auto* error = std::get_if(&message); + error && error->id) { + key = error->id->correlation_key(); + unauthorized = error->error.code == g_UNAUTHORIZED; + } + if (!key) { + return {}; + } + + std::lock_guard lock(owner->pending_mutex); + prune_expired_locked(*owner, Clock::now()); + const auto iter = owner->pending_requests.find(*key); + if (iter == owner->pending_requests.end()) { + return {}; + } + auto wire = unauthorized ? iter->second.wire : std::shared_ptr{}; + erase_pending_locked(*owner, iter); + return wire; + } catch (const std::exception&) { + return {}; + } + } + + struct ReadOperation { + explicit ReadOperation(std::shared_ptr state) : owner(std::move(state)) {} + + std::shared_ptr owner; + std::string raw; + std::shared_ptr retry_wire; + }; + + static Task read(std::shared_ptr owner) { + auto operation = std::make_shared(std::move(owner)); + return run_read(std::move(operation)); + } + + static Task run_read(std::shared_ptr operation) { + operation->raw = co_await operation->owner->inner->read_message(); + operation->retry_wire = take_retry_wire(operation->owner, operation->raw); + if (operation->retry_wire && co_await operation->owner->authenticator->try_refresh_token()) { + co_await write(operation->owner, operation->retry_wire); + operation->raw = co_await operation->owner->inner->read_message(); + (void)take_retry_wire(operation->owner, operation->raw); + } + co_return operation->raw; + } + + std::shared_ptr inner; + std::shared_ptr authenticator; + OAuthClientTransportOptions options; + std::mutex pending_mutex; + PendingMap pending_requests; + PendingOrder pending_order; + std::size_t pending_bytes{0}; + std::uint64_t next_generation{0}; + bool closed{false}; + bool uses_http_authorization_header{false}; +}; + +OAuthClientTransport::OAuthClientTransport(std::shared_ptr inner, + std::shared_ptr authenticator) + : OAuthClientTransport(std::move(inner), std::move(authenticator), {}) {} + +OAuthClientTransport::OAuthClientTransport(std::shared_ptr inner, + std::shared_ptr authenticator, + OAuthClientTransportOptions options) { + if (!inner || !authenticator) { + throw std::invalid_argument( + "OAuthClientTransport requires an inner transport and authenticator"); + } + + impl_ = std::make_shared(std::move(inner), std::move(authenticator), std::move(options)); + if (const auto http_transport = std::dynamic_pointer_cast(impl_->inner)) { + std::weak_ptr weak_authenticator = impl_->authenticator; + http_transport->set_bearer_token_provider([weak_authenticator]() { + const auto active_authenticator = weak_authenticator.lock(); + return active_authenticator ? active_authenticator->get_access_token() : std::string{}; + }); + impl_->uses_http_authorization_header = true; + } +} + +Task OAuthClientTransport::read_message() { return Impl::read(impl_); } + +Task OAuthClientTransport::write_message(std::string_view message) { + return Impl::write(impl_, std::make_shared(message)); +} + +void OAuthClientTransport::close() { + { + std::lock_guard lock(impl_->pending_mutex); + if (impl_->closed) { + return; + } + impl_->closed = true; + impl_->pending_requests.clear(); + impl_->pending_order.clear(); + impl_->pending_bytes = 0; + } + // Releases anything the authenticator has parked -- an in-flight discovery/token exchange, or a + // follower waiting on a single-flight timer -- before the inner transport is torn down, so a + // write() still working its way through the authorization retry loop fails promptly instead of + // outliving this call. + impl_->authenticator->close(); + impl_->inner->close(); +} + +} // namespace mcp::auth diff --git a/src/auth/oauth_internal.hpp b/src/auth/oauth_internal.hpp new file mode 100644 index 0000000..e49d2ce --- /dev/null +++ b/src/auth/oauth_internal.hpp @@ -0,0 +1,35 @@ +/** + * @file oauth_internal.hpp + * @brief Reach-in accessors for this SDK's own tests. NOT part of the shipped interface. + * + * This header lives under `src/` and is not installed, so nothing declared here reaches a consumer of + * the SDK. The symbols are exported because the test binary links the shared library, which is built + * with hidden visibility. Nothing in the library itself calls these. + */ +#pragma once + +#include +#include + +#include + +namespace mcp::auth::internal { + +/// How many per-scope abort records `client` is still holding on its own account. +/// +/// A client that remembers which scopes were aborted keeps one record per closed scope forever; +/// a client that keeps each latch on the scope itself has nothing to report and answers zero +/// however many scopes have come and gone. +MCP_API std::size_t retained_scope_record_count(const OAuthHttpClient& client); + +/// How many per-scope abort latches exist right now, across every client in this process. +/// +/// Counts latch objects, not references to them and not scopes. An in-flight exchange holds a +/// `shared_ptr` copy of an existing latch, so the count does NOT rise during a request, and a latch +/// outlives its scope while an exchange it issued is still running. Intended for leak checks: a count +/// that does not return to its starting value once the scopes are gone means something outlived its +/// scope. For a number that rises while a request is in flight, use `use_count()` on the shared +/// state. +MCP_API std::size_t live_scope_latch_count(); + +} // namespace mcp::auth::internal diff --git a/src/client/client.cpp b/src/client/client.cpp new file mode 100644 index 0000000..b89bc68 --- /dev/null +++ b/src/client/client.cpp @@ -0,0 +1,749 @@ +#include "../detail/diagnostic_text.hpp" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mcp { + +namespace { + +std::optional make_paginated_params(const std::optional& cursor) { + if (!cursor) { + return std::nullopt; + } + + PaginatedRequestParams params; + params.cursor = cursor; + return nlohmann::json(std::move(params)); +} + +} // namespace + +/// `error.message` is text the peer chose (dispatch_response() deserializes `error` with no +/// validation) and what() is what an application logs, so what() is flattened and bounded. error(), +/// and the message() it exposes, are structured protocol data a caller may compare or re-encode, and +/// stay exactly as the peer sent them; the split matches MetadataPolicyError's. +McpError::McpError(Error error) + : std::runtime_error("JSON-RPC error " + std::to_string(error.code) + ": " + + detail::sanitize_for_diagnostics(error.message)), + error_(std::move(error)) {} + +McpError::McpError(int code, std::string message, std::optional data) + : McpError(Error{code, std::move(message), std::move(data)}) {} + +struct Client::Impl { + struct PendingRequest { + PendingRequest(const boost::asio::any_io_executor& executor, std::chrono::milliseconds timeout) + : timer(executor), deadline(std::chrono::steady_clock::now() + timeout) { + timer.expires_at(deadline); + } + + boost::asio::steady_timer timer; + std::chrono::steady_clock::time_point deadline; + std::shared_ptr> cancel_write_before_start = + std::make_shared>(false); + nlohmann::json result; + std::optional error; + std::exception_ptr write_error; + bool completed{false}; + }; + + Impl(std::shared_ptr client_transport, const boost::asio::any_io_executor& executor, + ClientOptions client_options) + : transport(std::move(client_transport)), + strand(boost::asio::make_strand(executor)), + options(std::move(client_options)), + writer(transport, strand) {} + + static Task request(std::shared_ptr state, std::string_view method, + const std::optional& params, + const RequestOptions& request_options) { + const auto timeout = request_options.timeout.value_or(state->options.request_timeout); + if (timeout <= std::chrono::milliseconds::zero()) { + throw std::invalid_argument("Request timeout must be positive"); + } + + const int64_t id = state->next_request_id.fetch_add(1, std::memory_order_relaxed); + JSONRPCRequest request; + request.id = RequestId{std::to_string(id)}; + request.method = std::string(method); + request.params = params; + auto wire = std::make_shared(nlohmann::json(std::move(request)).dump()); + auto strand = state->strand; + // Awaiting post(strand) would not rebind a caller-owned coroutine's continuation. + // Spawn the complete pending-request lifecycle on the client strand instead. + return boost::asio::co_spawn(strand, + send_request_wire(std::move(state), std::move(wire), id, timeout), + boost::asio::use_awaitable); + } + + static Task notification(std::shared_ptr state, std::string_view method, + const std::optional& params) { + JSONRPCNotification notification; + notification.method = std::string(method); + notification.params = params; + auto wire = std::make_shared(nlohmann::json(std::move(notification)).dump()); + auto strand = state->strand; + // Keep the closed-state check and write initiation on the client strand. + return boost::asio::co_spawn(strand, send_notification_wire(std::move(state), std::move(wire)), + boost::asio::use_awaitable); + } + + static Task finish_connect(std::shared_ptr state, + Task initialize_request) { + try { + auto result_json = co_await std::move(initialize_request); + auto initialize_result = + std::make_shared(result_json.get()); + if (state->options.strict_protocol_validation && + !is_supported_protocol_version(initialize_result->protocolVersion)) { + request_close(state, "Server selected an unsupported protocol version"); + // The version is whatever the remote server put in its `initialize` result, and + // this message reaches the application's log. Flatten and bound it, or a server + // answering with CR/LF forges a log line the application did not write. + throw McpError(g_INVALID_REQUEST, "Server selected unsupported protocol version: " + + detail::sanitize_for_diagnostics( + initialize_result->protocolVersion)); + } + + co_await notification(state, "notifications/initialized", std::nullopt); + co_return std::move(*initialize_result); + } catch (...) { + request_close(state, "Client initialization failed"); + throw; + } + } + + static Task discard_result(Task request_task) { + static_cast(co_await std::move(request_task)); + } + + static void start_read_loop(const std::shared_ptr& state) { + boost::asio::co_spawn(state->strand, read_loop(state), boost::asio::detached); + } + + static void request_close(const std::shared_ptr& state, std::string message) { + if (state->closed.exchange(true, std::memory_order_acq_rel)) { + return; + } + + boost::asio::post(state->strand, [state, message = std::move(message)]() mutable { + fail_pending_requests(state, Error{g_CONNECTION_CLOSED, std::move(message), std::nullopt}); + state->transport->close(); + }); + } + + static Task send_request_wire(std::shared_ptr state, + std::shared_ptr wire, int64_t id, + std::chrono::milliseconds timeout) { + if (state->closed.load(std::memory_order_acquire)) { + throw McpError(g_CONNECTION_CLOSED, "Client transport is closed"); + } + + const auto id_key = RequestId{std::to_string(id)}.correlation_key(); + PendingRequest* pending_request = nullptr; + { + std::lock_guard lock(state->pending_requests_mutex); + auto [iter, inserted] = state->pending_requests.try_emplace( + id_key, std::make_unique(state->strand, timeout)); + if (!inserted) { + throw McpError(g_INVALID_REQUEST, "Duplicate pending request id: " + id_key); + } + pending_request = iter->second.get(); + } + + boost::asio::co_spawn( + state->strand, + state->writer.write_message(std::move(wire), pending_request->cancel_write_before_start), + [state, id_key](std::exception_ptr write_error) { + complete_request_write(state, id_key, std::move(write_error)); + }); + + try { + co_await pending_request->timer.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + erase_pending_request(state, id_key); + throw; + } + } + + nlohmann::json result; + std::optional error; + std::exception_ptr write_error; + { + std::lock_guard lock(state->pending_requests_mutex); + auto iter = state->pending_requests.find(id_key); + if (iter == state->pending_requests.end()) { + throw McpError(g_CONNECTION_CLOSED, "Pending request was removed: " + id_key); + } + if (!iter->second->completed) { + iter->second->cancel_write_before_start->store(true, std::memory_order_release); + iter->second->error = Error{g_REQUEST_TIMEOUT, "Request timed out", + nlohmann::json{{"id", std::to_string(id)}}}; + iter->second->completed = true; + } + result = std::move(iter->second->result); + error = std::move(iter->second->error); + write_error = std::move(iter->second->write_error); + state->pending_requests.erase(iter); + } + + if (error) { + throw McpError(std::move(*error)); + } + if (write_error) { + std::rethrow_exception(write_error); + } + co_return result; + } + + static Task send_notification_wire(std::shared_ptr state, + std::shared_ptr wire) { + if (state->closed.load(std::memory_order_acquire)) { + throw McpError(g_CONNECTION_CLOSED, "Client transport is closed"); + } + + try { + co_await state->writer.write_message(wire); + } catch (...) { + if (state->closed.load(std::memory_order_acquire)) { + throw McpError(g_CONNECTION_CLOSED, "Client transport is closed"); + } + throw; + } + } + + /// Two failure classes meet in this loop and only one is fatal. read_message() failing means the + /// transport can no longer produce messages, so the session ends: pending requests fail and the + /// transport closes; it is the only statement inside the outer try for that reason. A failure + /// after it concerns one message already off the wire: it is reported through on_protocol_error + /// and dropped, and the loop continues. + static Task read_loop(std::shared_ptr state) { + try { + for (;;) { + auto raw = co_await state->transport->read_message(); + dispatch_message(state, raw); + } + } catch (const std::exception& error) { + state->closed.store(true, std::memory_order_release); + fail_pending_requests(state, Error{g_CONNECTION_CLOSED, + "Client read loop stopped: " + + detail::sanitize_for_diagnostics(error.what()), + std::nullopt}); + state->transport->close(); + } catch (...) { + state->closed.store(true, std::memory_order_release); + fail_pending_requests(state, + Error{g_CONNECTION_CLOSED, "Client read loop stopped", std::nullopt}); + state->transport->close(); + } + } + + /// Decodes and routes one received message. Never throws: an unusable message is reported + /// and dropped so that a single bad message cannot end the session. + static void dispatch_message(const std::shared_ptr& state, const std::string& raw) { + try { + auto json_message = nlohmann::json::parse(raw); + + if (!json_message.is_object()) { + throw McpError(g_INVALID_REQUEST, "JSON-RPC message must be an object"); + } + if (state->options.strict_protocol_validation) { + if (!json_message.contains("jsonrpc") || !json_message.at("jsonrpc").is_string()) { + throw McpError(g_INVALID_REQUEST, "JSON-RPC message is missing jsonrpc"); + } + detail::validate_jsonrpc_version(json_message.at("jsonrpc").get()); + } + + const bool has_id = json_message.contains("id"); + const bool has_method = json_message.contains("method"); + if (has_method && !json_message.at("method").is_string()) { + throw McpError(g_INVALID_REQUEST, "JSON-RPC method must be a string"); + } + + if (has_id && !has_method) { + dispatch_response(state, json_message); + } else if (has_id && has_method) { + // dispatch_incoming_request is detached, so anything that throws inside it is + // swallowed and the peer never hears back. Reject an unusable id here, where the + // drop is reported, instead of leaving the request silently unanswered. + const auto& id = json_message.at("id"); + if (!id.is_string() && !id.is_number_integer()) { + throw McpError(g_INVALID_REQUEST, + "JSON-RPC request id must be a string or integer"); + } + boost::asio::co_spawn(state->strand, + dispatch_incoming_request(state, std::move(json_message)), + boost::asio::detached); + } else if (!has_id && has_method) { + dispatch_notification(state, json_message); + } + } catch (const std::exception& error) { + report_protocol_error(state, g_PARSE_ERROR, + "Dropped an undecodable message from the peer: " + + detail::sanitize_for_diagnostics(error.what())); + } catch (...) { + report_protocol_error(state, g_PARSE_ERROR, "Dropped an undecodable message from the peer"); + } + } + + static void report_protocol_error(const std::shared_ptr& state, int code, + std::string message) { + if (!state->options.on_protocol_error) { + return; + } + try { + state->options.on_protocol_error(Error{code, std::move(message), std::nullopt}); + } catch (...) { + // The hook exists to report failures; a failure of the hook itself has no listener. + } + } + + static void dispatch_response(const std::shared_ptr& state, + const nlohmann::json& json_message) { + const auto& id = json_message.at("id"); + if (!id.is_string() && !id.is_number_integer()) { + return; + } + const auto id_key = id.get().correlation_key(); + + std::lock_guard lock(state->pending_requests_mutex); + auto iter = state->pending_requests.find(id_key); + if (iter == state->pending_requests.end()) { + return; + } + if (std::chrono::steady_clock::now() >= iter->second->deadline) { + return; + } + + // Asymmetric on purpose: a null `result` is a legitimate empty result, so presence is the + // right test for it, while a null `error` is the absence of an error and must not count. + const bool has_error = detail::has_json_value(json_message, "error"); + const bool has_result = json_message.contains("result"); + if (has_error == has_result) { + iter->second->error = + Error{g_INVALID_REQUEST, + "JSON-RPC response must contain exactly one of result or error", std::nullopt}; + } else if (has_error) { + // An undecodable `error` member still identifies the request it answers. Fail that + // request now rather than drop the response and leave the caller waiting out its + // deadline for an answer that has already arrived. + try { + iter->second->error = json_message.at("error").get(); + } catch (const std::exception& decode_error) { + iter->second->error = Error{g_INVALID_REQUEST, + "Malformed JSON-RPC error object: " + + detail::sanitize_for_diagnostics(decode_error.what()), + std::nullopt}; + } + } else { + iter->second->result = json_message.at("result"); + } + + iter->second->completed = true; + signal_pending_request(*iter->second); + } + + static void dispatch_notification(const std::shared_ptr& state, + const nlohmann::json& json_message) { + const auto method = json_message.at("method").get(); + if (method == "notifications/cancelled") { + if (detail::has_json_value(json_message, "params")) { + auto params = json_message.at("params").get(); + const auto id_key = params.requestId.correlation_key(); + + std::lock_guard lock(state->pending_requests_mutex); + auto iter = state->pending_requests.find(id_key); + if (iter != state->pending_requests.end() && + std::chrono::steady_clock::now() < iter->second->deadline) { + iter->second->error = Error{g_REQUEST_CANCELLED, "Request cancelled by server"}; + iter->second->completed = true; + signal_pending_request(*iter->second); + } + } + return; + } + + NotificationCallback callback; + { + std::lock_guard lock(state->handlers_mutex); + auto iter = state->notification_handlers.find(method); + if (iter != state->notification_handlers.end()) { + callback = iter->second; + } + } + if (callback) { + auto params = detail::has_json_value(json_message, "params") ? json_message.at("params") + : nlohmann::json::object(); + // The callback is application code running on the read loop. Letting it throw past + // here would end the session over a bug the application could otherwise handle, so + // the failure is reported against its method and the session carries on. + try { + callback(params); + } catch (const std::exception& error) { + report_protocol_error(state, g_INTERNAL_ERROR, + "Notification callback for " + + detail::sanitize_for_diagnostics(method) + + " threw: " + detail::sanitize_for_diagnostics(error.what())); + } catch (...) { + report_protocol_error( + state, g_INTERNAL_ERROR, + "Notification callback for " + detail::sanitize_for_diagnostics(method) + " threw"); + } + } + } + + static Task dispatch_incoming_request(std::shared_ptr state, + nlohmann::json json_message) { + std::string_view method = json_message.at("method").get_ref(); + + RequestHandler request_handler; + { + std::lock_guard lock(state->handlers_mutex); + auto iter = state->request_handlers.find(method); + if (iter != state->request_handlers.end()) { + request_handler = iter->second; + } + } + + if (request_handler) { + auto params = detail::has_json_value(json_message, "params") ? json_message.at("params") + : nlohmann::json::object(); + nlohmann::json error_payload; + nlohmann::json result; + try { + result = co_await request_handler(params); + } catch (const std::exception& error) { + error_payload = error.what(); + } catch (...) { + error_payload = "Request handler failed"; + } + + if (!error_payload.is_null()) { + co_await state->writer.write_message( + make_error_wire(json_message.at("id").get(), g_INTERNAL_ERROR, + error_payload.get())); + } else { + co_await state->writer.write_message( + make_result_wire(json_message.at("id").get(), std::move(result))); + } + } else if (method == "ping") { + co_await state->writer.write_message( + make_result_wire(json_message.at("id").get(), nlohmann::json::object())); + } else { + co_await state->writer.write_message( + make_error_wire(json_message.at("id").get(), g_METHOD_NOT_FOUND, + "Method not found: " + std::string(method))); + } + } + + static std::string make_result_wire(const RequestId& id, nlohmann::json result) { + JSONRPCResultResponse response; + response.id = id; + response.result = std::move(result); + return nlohmann::json(std::move(response)).dump(); + } + + static std::string make_error_wire(const RequestId& id, int code, std::string message) { + Error error; + error.code = code; + error.message = std::move(message); + JSONRPCErrorResponse response; + response.id = id; + response.error = std::move(error); + return nlohmann::json(std::move(response)).dump(); + } + + static void signal_pending_request(PendingRequest& pending_request) { + pending_request.timer.expires_at(std::chrono::steady_clock::now()); + } + + static void complete_request_write(const std::shared_ptr& state, const std::string& id_key, + std::exception_ptr write_error) { + if (!write_error) { + return; + } + + std::lock_guard lock(state->pending_requests_mutex); + auto iter = state->pending_requests.find(id_key); + if (iter == state->pending_requests.end() || iter->second->completed || + std::chrono::steady_clock::now() >= iter->second->deadline) { + return; + } + + if (state->closed.load(std::memory_order_acquire)) { + iter->second->error = Error{g_CONNECTION_CLOSED, "Client transport is closed"}; + } else { + iter->second->write_error = std::move(write_error); + } + iter->second->completed = true; + signal_pending_request(*iter->second); + } + + static void erase_pending_request(const std::shared_ptr& state, const std::string& id_key) { + std::lock_guard lock(state->pending_requests_mutex); + state->pending_requests.erase(id_key); + } + + static void fail_pending_requests(const std::shared_ptr& state, const Error& error) { + std::lock_guard lock(state->pending_requests_mutex); + for (auto& [id, pending_request] : state->pending_requests) { + static_cast(id); + if (!pending_request->completed) { + pending_request->error = error; + pending_request->completed = true; + } + signal_pending_request(*pending_request); + } + } + + std::shared_ptr transport; + boost::asio::strand strand; + ClientOptions options; + detail::SerializedTransportWriter writer; + + mutable std::mutex pending_requests_mutex; + std::map, std::less<>> pending_requests; + std::atomic next_request_id{1}; + std::atomic read_loop_started{false}; + std::atomic closed{false}; + + mutable std::mutex handlers_mutex; + std::map> notification_handlers; + std::map> request_handlers; + std::vector roots; + ClientCapabilities client_capabilities; +}; + +Client::Client(std::shared_ptr transport, const boost::asio::any_io_executor& executor, + ClientOptions options) { + if (!transport) { + throw std::invalid_argument("Client transport must not be null"); + } + if (options.request_timeout <= std::chrono::milliseconds::zero()) { + throw std::invalid_argument("Client request timeout must be positive"); + } + impl_ = std::make_shared(std::move(transport), executor, std::move(options)); +} + +Client::~Client() { close(); } + +Task Client::connect(std::string_view name, std::string_view version) { + Implementation info; + info.name = std::string(name); + info.version = std::string(version); + return connect(info, {}); +} + +Task Client::connect(const Implementation& client_info, + const ClientCapabilities& capabilities) { + auto state = impl_; + if (state->closed.load(std::memory_order_acquire)) { + throw McpError(g_CONNECTION_CLOSED, "Client is closed"); + } + if (state->read_loop_started.exchange(true, std::memory_order_acq_rel)) { + throw McpError(g_INVALID_REQUEST, "Client has already been initialized"); + } + + Impl::start_read_loop(state); + InitializeRequest initialize_request; + initialize_request.protocolVersion = std::string(g_LATEST_PROTOCOL_VERSION); + initialize_request.clientInfo = client_info; + { + std::lock_guard lock(state->handlers_mutex); + state->client_capabilities = capabilities; + if (state->request_handlers.contains("roots/list") && + !state->client_capabilities.roots.has_value()) { + state->client_capabilities.roots = ClientCapabilities::RootsCapability{}; + } + initialize_request.capabilities = state->client_capabilities; + } + auto request_task = + Impl::request(state, "initialize", nlohmann::json(std::move(initialize_request)), {}); + return Impl::finish_connect(std::move(state), std::move(request_task)); +} + +Task Client::send_request(std::string_view method, + const std::optional& params) { + return send_request(method, params, {}); +} + +Task Client::send_request(std::string_view method, + const std::optional& params, + const RequestOptions& request_options) { + if (!impl_->read_loop_started.load(std::memory_order_acquire)) { + throw McpError(g_INVALID_REQUEST, "Client must connect before sending requests"); + } + return Impl::request(impl_, method, params, request_options); +} + +Task Client::send_notification(std::string_view method, + const std::optional& params) { + if (!impl_->read_loop_started.load(std::memory_order_acquire)) { + throw McpError(g_INVALID_REQUEST, "Client must connect before sending notifications"); + } + return Impl::notification(impl_, method, params); +} + +std::size_t Client::pending_request_count() const { + std::lock_guard lock(impl_->pending_requests_mutex); + return impl_->pending_requests.size(); +} + +void Client::close() { + if (impl_) { + Impl::request_close(impl_, "Client transport closed"); + } +} + +Task Client::list_tools(const std::optional& cursor) { + return call_and_parse("tools/list", make_paginated_params(cursor)); +} + +Task Client::list_resources(const std::optional& cursor) { + return call_and_parse("resources/list", make_paginated_params(cursor)); +} + +Task Client::read_resource(const std::string& uri) { + ReadResourceRequestParams params; + params.uri = uri; + return call_and_parse("resources/read", nlohmann::json(std::move(params))); +} + +Task Client::list_resource_templates( + const std::optional& cursor) { + return call_and_parse("resources/templates/list", + make_paginated_params(cursor)); +} + +Task Client::list_prompts(const std::optional& cursor) { + return call_and_parse("prompts/list", make_paginated_params(cursor)); +} + +Task Client::get_prompt( + const std::string& name, const std::optional>& arguments) { + GetPromptRequestParams params; + params.name = name; + params.arguments = arguments; + return call_and_parse("prompts/get", nlohmann::json(std::move(params))); +} + +Task Client::complete(const CompleteParams& params) { + return call_and_parse("completion/complete", nlohmann::json(params)); +} + +Task Client::ping() { + if (!impl_->read_loop_started.load(std::memory_order_acquire)) { + throw McpError(g_INVALID_REQUEST, "Client must connect before sending requests"); + } + return Impl::discard_result(Impl::request(impl_, "ping", std::nullopt, {})); +} + +Task Client::cancel(const RequestId& request_id, const std::optional& reason) { + CancelledNotificationParams params; + params.requestId = request_id; + params.reason = reason; + return send_notification("notifications/cancelled", nlohmann::json(std::move(params))); +} + +void Client::on_notification(const std::string& method, NotificationCallback callback) { + std::lock_guard lock(impl_->handlers_mutex); + impl_->notification_handlers[method] = std::move(callback); +} + +void Client::on_progress(ProgressCallback callback) { + on_notification( + "notifications/progress", [callback = std::move(callback)](const nlohmann::json& params) { + ProgressNotificationParams progress; + try { + params.get_to(progress); + } catch (const std::exception& error) { + // Reported through ClientOptions::on_protocol_error by the read loop. + // Naming the cause here separates a notification the peer sent wrong + // from a bug in the application's own progress callback. + throw McpError(g_INVALID_PARAMS, "Malformed progress notification params: " + + detail::sanitize_for_diagnostics(error.what())); + } + callback(progress); + }); +} + +void Client::on_request(const std::string& method, RequestHandler handler) { + std::lock_guard lock(impl_->handlers_mutex); + impl_->request_handlers[method] = std::move(handler); +} + +void Client::on_elicitation(std::function(const ElicitRequestParams&)> handler) { + on_request("elicitation/create", + [handler = std::move(handler)](const nlohmann::json& params) -> Task { + auto result = co_await handler(params.get()); + co_return nlohmann::json(std::move(result)); + }); +} + +void Client::set_roots(const std::vector& roots, bool notify) { + auto state = impl_; + bool should_notify = false; + { + std::lock_guard lock(state->handlers_mutex); + state->roots = roots; + if (!state->request_handlers.contains("roots/list")) { + std::weak_ptr weak_state = state; + state->request_handlers.emplace( + "roots/list", [weak_state](const nlohmann::json&) -> Task { + auto active_state = weak_state.lock(); + if (!active_state) { + throw McpError(g_CONNECTION_CLOSED, "Client is closed"); + } + + ListRootsResult result; + { + std::lock_guard roots_lock(active_state->handlers_mutex); + result.roots = active_state->roots; + } + co_return nlohmann::json(std::move(result)); + }); + } + + should_notify = notify && state->client_capabilities.roots.has_value() && + state->client_capabilities.roots->listChanged.value_or(false); + } + + if (should_notify && !state->closed.load(std::memory_order_acquire)) { + boost::asio::co_spawn( + state->strand, Impl::notification(state, "notifications/roots/list_changed", std::nullopt), + boost::asio::detached); + } +} + +void Client::on_roots_list(std::function(const nlohmann::json&)> handler) { + on_request("roots/list", + [handler = std::move(handler)](const nlohmann::json& params) -> Task { + auto result = co_await handler(params); + co_return nlohmann::json(std::move(result)); + }); +} + +} // namespace mcp diff --git a/src/core/secure_random.cpp b/src/core/secure_random.cpp new file mode 100644 index 0000000..3508969 --- /dev/null +++ b/src/core/secure_random.cpp @@ -0,0 +1,34 @@ +#include +#include + +#include + +#include +#include +#include + +namespace mcp::detail { + +std::string generate_secure_session_id() { + const std::size_t byte_count = (constants::g_session_id_length + 1) / 2; + std::vector random_bytes(byte_count); + if (RAND_bytes(random_bytes.data(), static_cast(random_bytes.size())) != 1) { + throw std::runtime_error("Failed to generate a cryptographically secure session id"); + } + + std::string session_id; + session_id.reserve(constants::g_session_id_length); + for (const unsigned char value : random_bytes) { + session_id.push_back(constants::g_hex_digits[(value >> 4U) & 0x0FU]); + if (session_id.size() == constants::g_session_id_length) { + break; + } + session_id.push_back(constants::g_hex_digits[value & 0x0FU]); + if (session_id.size() == constants::g_session_id_length) { + break; + } + } + return session_id; +} + +} // namespace mcp::detail diff --git a/src/core/serialized_transport_writer.cpp b/src/core/serialized_transport_writer.cpp new file mode 100644 index 0000000..4c5dcbe --- /dev/null +++ b/src/core/serialized_transport_writer.cpp @@ -0,0 +1,168 @@ +#include + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace mcp::detail { + +struct SerializedTransportWriterState { + struct Operation { + Operation(const boost::asio::any_io_executor& executor, + std::shared_ptr owned_message, + std::shared_ptr> cancellation) + : message(std::move(owned_message)), + cancel_before_start(std::move(cancellation)), + completion(executor) { + completion.expires_at(std::chrono::steady_clock::time_point::max()); + } + + std::shared_ptr message; + std::shared_ptr> cancel_before_start; + boost::asio::steady_timer completion; + std::exception_ptr error; + bool completed{false}; + }; + + SerializedTransportWriterState(std::shared_ptr owned_transport, + const boost::asio::any_io_executor& executor) + : transport(std::move(owned_transport)), strand(boost::asio::make_strand(executor)) {} + + static void complete(const std::shared_ptr& operation, + std::exception_ptr error = nullptr) { + operation->error = std::move(error); + operation->completed = true; + + boost::system::error_code ignored; + operation->completion.expires_at(std::chrono::steady_clock::time_point::min(), ignored); + } + + void fail_queue(std::exception_ptr error) { + failure = error; + while (!queue.empty()) { + auto operation = std::move(queue.front()); + queue.pop_front(); + complete(operation, error); + } + draining = false; + } + + static Task drain(std::shared_ptr state) { + while (!state->queue.empty()) { + auto operation = state->queue.front(); + if (operation->cancel_before_start && + operation->cancel_before_start->load(std::memory_order_acquire)) { + state->queue.pop_front(); + complete(operation, std::make_exception_ptr(boost::system::system_error( + boost::asio::error::operation_aborted))); + continue; + } + try { + co_await state->transport->write_message(*operation->message); + } catch (...) { + state->fail_queue(std::current_exception()); + co_return; + } + + state->queue.pop_front(); + complete(operation); + } + + state->draining = false; + } + + static Task enqueue_and_wait(std::shared_ptr state, + std::shared_ptr message, + std::shared_ptr> cancel_before_start) { + if (state->failure) { + std::rethrow_exception(state->failure); + } + + auto operation = std::make_shared(state->strand, std::move(message), + std::move(cancel_before_start)); + state->queue.push_back(operation); + + if (!state->draining) { + state->draining = true; + boost::asio::co_spawn(state->strand, drain(state), boost::asio::detached); + } + + while (!operation->completed) { + boost::system::error_code ignored; + co_await operation->completion.async_wait( + boost::asio::redirect_error(boost::asio::use_awaitable, ignored)); + } + + if (operation->error) { + std::rethrow_exception(operation->error); + } + } + + std::shared_ptr transport; + boost::asio::strand strand; + std::deque> queue; + std::exception_ptr failure; + bool draining{false}; +}; + +SerializedTransportWriter::SerializedTransportWriter(std::shared_ptr transport, + const boost::asio::any_io_executor& executor) + : state_(std::make_shared(std::move(transport), executor)) { + if (!state_->transport) { + throw std::invalid_argument("SerializedTransportWriter requires a transport"); + } +} + +Task SerializedTransportWriter::write_message(std::string_view message) const { + if (!state_) { + throw std::logic_error("Cannot use a moved-from SerializedTransportWriter"); + } + + // This function deliberately is not a coroutine. Copying here makes the + // string_view safe before the returned awaitable can suspend, while the + // shared allocation avoids keeping an SSO string in a GCC 11 coroutine frame. + auto owned_message = std::make_shared(message); + auto state = state_; + return boost::asio::co_spawn( + state->strand, + SerializedTransportWriterState::enqueue_and_wait(state, std::move(owned_message), nullptr), + boost::asio::use_awaitable); +} + +Task SerializedTransportWriter::write_message(std::shared_ptr message) const { + return write_message(std::move(message), nullptr); +} + +Task SerializedTransportWriter::write_message( + std::shared_ptr message, + std::shared_ptr> cancel_before_start) const { + if (!state_) { + throw std::logic_error("Cannot use a moved-from SerializedTransportWriter"); + } + if (!message) { + throw std::invalid_argument("SerializedTransportWriter message must not be null"); + } + auto state = state_; + return boost::asio::co_spawn(state->strand, + SerializedTransportWriterState::enqueue_and_wait( + state, std::move(message), std::move(cancel_before_start)), + boost::asio::use_awaitable); +} + +} // namespace mcp::detail diff --git a/src/detail/diagnostic_text.hpp b/src/detail/diagnostic_text.hpp new file mode 100644 index 0000000..8500247 --- /dev/null +++ b/src/detail/diagnostic_text.hpp @@ -0,0 +1,21 @@ +#pragma once + +#include + +#include +#include + +namespace mcp::detail { + +/// Flattens and bounds peer-controlled text for a diagnostic, for throw sites outside `src/auth/`. +/// +/// A plain forwarder to `mcp::auth::detail::sanitize_for_diagnostics`; see its declaration in +/// `mcp/auth/metadata_policy.hpp` for exactly what "flatten" and "bound" mean. That declaration is +/// `MCP_API`-exported, so it cannot be moved or renamed without breaking callers. This header exists +/// so `src/client/` and `src/server/` do not each grow their own `#include ` for their +/// only dependency on the auth headers. +[[nodiscard]] inline std::string sanitize_for_diagnostics(std::string_view value) { + return auth::detail::sanitize_for_diagnostics(value); +} + +} // namespace mcp::detail diff --git a/src/protocol/tools.cpp b/src/protocol/tools.cpp new file mode 100644 index 0000000..1b54fd8 --- /dev/null +++ b/src/protocol/tools.cpp @@ -0,0 +1,28 @@ +#include + +#include + +namespace mcp { + +CallToolResult make_tool_text_result(std::string text) { + CallToolResult result; + TextContent content; + content.text = std::move(text); + result.content.emplace_back(std::move(content)); + return result; +} + +CallToolResult make_tool_error_result(std::string message) { + auto result = make_tool_text_result(std::move(message)); + result.isError = true; + return result; +} + +CallToolResult make_tool_structured_result(nlohmann::json structured_content, + std::optional text) { + auto result = make_tool_text_result(text ? std::move(*text) : structured_content.dump()); + result.structuredContent = std::move(structured_content); + return result; +} + +} // namespace mcp diff --git a/src/server/server.cpp b/src/server/server.cpp index d72a7e9..6410dc2 100644 --- a/src/server/server.cpp +++ b/src/server/server.cpp @@ -1,3 +1,6 @@ +#include "../detail/diagnostic_text.hpp" + +#include #include #include @@ -9,19 +12,28 @@ // [scope-before-await] Build SSO-risky objects in {}, serialise to wire string, then co_await. // [wire-builders] make_result_wire/make_error_wire are synchronous helpers; // do NOT convert them to Task coroutines. +// +// These conventions are the only guard. GCC 11 is also sensitive to the shape of the coroutine +// frames, so an unrelated refactor can move a failure in or out of existence; a green ubuntu-22.04 +// means the bug is not currently tripped, not that it is fixed. #include #include +#include #include #include #include +#include #include +#include #include #include #include +#include #include #include +#include #include #include #include @@ -31,19 +43,295 @@ namespace mcp { +namespace { + +struct UriTemplatePattern { + std::string source; + std::regex matcher; +}; + +// The input is bounded because the matcher cannot be trusted with it, not as a policy on resource +// names. Every std::regex implementation spends stack in proportion to the subject's length, by an +// amount each decides for itself, and the smallest default thread stack this SDK runs on -- 512 KB +// for a non-main thread on macOS -- has to survive the longest URI accepted here. Matching an RFC +// 6570 level-2 template ({var}, {+var}) needs no backtracking, so a linear, stack-free segment +// matcher would retire this limit entirely. +constexpr std::size_t g_MAX_TEMPLATE_MATCH_URI_LENGTH = 512; + +bool is_uri_template_operator(char ch) { + constexpr std::string_view operators = "+#./;?&"; + return operators.find(ch) != std::string_view::npos; +} + +void validate_uri_template_variables(std::string_view expression) { + if (!expression.empty() && is_uri_template_operator(expression.front())) { + expression.remove_prefix(1); + } + if (expression.empty()) { + throw std::invalid_argument("Resource URI template contains an empty expression"); + } + + while (!expression.empty()) { + auto separator = expression.find(','); + auto variable = expression.substr(0, separator); + if (variable.empty()) { + throw std::invalid_argument("Resource URI template contains an empty variable"); + } + + auto modifier = variable.find_first_of(":*"); + auto name = variable.substr(0, modifier); + if (name.empty() || !std::ranges::all_of(name, [](char ch) { + auto uch = static_cast(ch); + return std::isalnum(uch) != 0 || ch == '_' || ch == '.' || ch == '%'; + })) { + throw std::invalid_argument("Resource URI template contains an invalid variable name"); + } + + if (separator == std::string_view::npos) { + break; + } + expression.remove_prefix(separator + 1); + } +} + +void append_regex_literal(std::string& pattern, char ch) { + constexpr std::string_view metacharacters = R"(\.^$|()[]*+?{})"; + if (metacharacters.find(ch) != std::string_view::npos) { + pattern.push_back('\\'); + } + pattern.push_back(ch); +} + +void append_uri_expression_pattern(std::string& pattern, std::string_view expression) { + validate_uri_template_variables(expression); + char expression_operator = is_uri_template_operator(expression.front()) ? expression.front() : 0; + + switch (expression_operator) { + case '+': + pattern += R"([^?#]+)"; + break; + case '#': + pattern += R"((?:#[^#]*)?)"; + break; + case '.': + pattern += R"((?:\.[^/?#]+(?:\.[^/?#]+)*)?)"; + break; + case '/': + pattern += R"((?:/[^?#]+)?)"; + break; + case ';': + pattern += R"((?:;[^?#]*)?)"; + break; + case '?': + pattern += R"((?:\?[^#]*)?)"; + break; + case '&': + pattern += R"((?:&[^#]*)?)"; + break; + default: + pattern += R"([^/?#]+)"; + break; + } +} + +UriTemplatePattern compile_uri_template(std::string_view uri_template) { + if (uri_template.empty()) { + throw std::invalid_argument("Resource URI template must not be empty"); + } + + std::string pattern = "^"; + bool previous_token_was_expression = false; + for (std::size_t index = 0; index < uri_template.size();) { + if (uri_template[index] == '}') { + throw std::invalid_argument("Resource URI template contains an unmatched '}'"); + } + if (uri_template[index] != '{') { + append_regex_literal(pattern, uri_template[index]); + previous_token_was_expression = false; + ++index; + continue; + } + + auto close = uri_template.find('}', index + 1); + if (close == std::string_view::npos) { + throw std::invalid_argument("Resource URI template contains an unmatched '{'"); + } + // Two expressions with nothing between them compile to adjacent unbounded runs, which the + // regex engine can only resolve by backtracking over every way of splitting the input. The + // boundary between such variables is undecidable anyway, so reject the template outright. + if (previous_token_was_expression) { + throw std::invalid_argument( + "Resource URI template must separate adjacent expressions with a literal"); + } + auto expression = uri_template.substr(index + 1, close - index - 1); + append_uri_expression_pattern(pattern, expression); + previous_token_was_expression = true; + index = close + 1; + } + pattern += '$'; + + return UriTemplatePattern{pattern, std::regex(pattern, std::regex::ECMAScript)}; +} + +void validate_call_tool_result(const nlohmann::json& result) { + if (!result.is_object() || !result.contains("content")) { + throw std::invalid_argument("Raw tool result must be a CallToolResult object with content"); + } + if (result.contains("structuredContent") && !result.at("structuredContent").is_object()) { + throw std::invalid_argument("CallToolResult structuredContent must be a JSON object"); + } + if (result.contains("_meta") && !result.at("_meta").is_object()) { + throw std::invalid_argument("CallToolResult _meta must be a JSON object"); + } + + static_cast(result.get()); +} + +bool is_serialized_call_tool_result(const nlohmann::json& result) { + try { + validate_call_tool_result(result); + } catch (const std::exception&) { + return false; + } + return true; +} + +nlohmann::json normalize_structured_tool_result(nlohmann::json result) { + if (result.is_object()) { + // A handler that built a CallToolResult is already speaking the protocol; wrapping it again + // would bury its content blocks inside a text block. Domain data that merely carries a + // "content" key fails the full CallToolResult check and is still wrapped. + if (is_serialized_call_tool_result(result)) { + return result; + } + return nlohmann::json(make_tool_structured_result(std::move(result))); + } + if (result.is_string()) { + return nlohmann::json(make_tool_text_result(result.get())); + } + return nlohmann::json(make_tool_text_result(result.dump())); +} + +class InvalidParamsError final : public std::invalid_argument { + public: + using std::invalid_argument::invalid_argument; +}; + +template +Params deserialize_request_params(const nlohmann::json& message, std::string_view method) { + try { + return message.at("params").get(); + } catch (const nlohmann::json::exception& error) { + throw InvalidParamsError("Invalid " + std::string(method) + " params: " + error.what()); + } catch (const std::invalid_argument& error) { + throw InvalidParamsError("Invalid " + std::string(method) + " params: " + error.what()); + } +} + +bool has_valid_request_id(const nlohmann::json& message) { + return message.contains("id") && + (message.at("id").is_string() || message.at("id").is_number_integer()); +} + +// A peer that serializes an absent optional as an explicit null means the member is not there, and +// a null is never a usable error object. "result" gets no such treatment: JSON-RPC allows a null +// result as a legitimate empty value, so a present-but-null result is a result. +bool has_error_member(const nlohmann::json& message) { + return message.contains("error") && !message.at("error").is_null(); +} + +const char* validate_request_envelope(const nlohmann::json& message) { + if (!message.is_object()) { + return "JSON-RPC request must be an object"; + } + if (!message.contains("jsonrpc") || !message.at("jsonrpc").is_string() || + message.at("jsonrpc") != "2.0") { + return "jsonrpc must be \"2.0\""; + } + if (!has_valid_request_id(message)) { + return "JSON-RPC request id must be a string or integer"; + } + if (!message.contains("method") || !message.at("method").is_string()) { + return "JSON-RPC request method must be a string"; + } + if (message.contains("result") || has_error_member(message)) { + return "JSON-RPC request must not contain result or error"; + } + if (message.contains("params") && !message.at("params").is_object()) { + return "MCP request params must be an object"; + } + return nullptr; +} + +bool is_valid_notification_envelope(const nlohmann::json& message) { + // A notification is never answered, so dropping one costs the peer a silent protocol stall + // rather than an error. An explicit null for an absent optional is the default shape for + // several mainstream JSON serializers, so it is read as "no params" rather than rejected. + return message.is_object() && !message.contains("id") && message.contains("jsonrpc") && + message.at("jsonrpc").is_string() && message.at("jsonrpc") == "2.0" && + message.contains("method") && message.at("method").is_string() && + !message.contains("result") && !has_error_member(message) && + (!message.contains("params") || message.at("params").is_null() || + message.at("params").is_object()); +} + +bool is_valid_response_envelope(const nlohmann::json& message) { + if (!message.is_object() || !has_valid_request_id(message) || message.contains("method") || + !message.contains("jsonrpc") || !message.at("jsonrpc").is_string() || + message.at("jsonrpc") != "2.0" || (message.contains("result") == has_error_member(message))) { + return false; + } + + if (!has_error_member(message)) { + return true; + } + + try { + static_cast(message.at("error").get()); + return true; + } catch (const std::exception&) { + return false; + } +} + +std::string make_invalid_request_wire(const nlohmann::json& message, std::string_view reason) { + nlohmann::json id = nullptr; + if (message.is_object() && has_valid_request_id(message)) { + id = message.at("id"); + } + return nlohmann::json{{"jsonrpc", "2.0"}, + {"id", std::move(id)}, + {"error", {{"code", g_INVALID_REQUEST}, {"message", reason}}}} + .dump(); +} + +std::string make_parse_error_wire() { + return nlohmann::json{{"jsonrpc", "2.0"}, + {"id", nullptr}, + {"error", {{"code", g_PARSE_ERROR}, {"message", "Parse error"}}}} + .dump(); +} + +} // namespace + struct Server::PendingRequest { - std::unique_ptr timer; + std::shared_ptr timer; nlohmann::json result; std::optional error; + bool completed{false}; }; struct Server::Session { std::shared_ptr transport; MemoryTransport* memory_transport = nullptr; std::unique_ptr> strand; + std::shared_ptr writer; std::map pending_requests; std::map>> in_flight; std::map subscriptions; + std::shared_ptr drain_timer; + std::atomic_size_t active_dispatches{0}; + std::atomic_bool stopping{false}; }; class NullTransport : public ITransport { @@ -62,12 +350,37 @@ NullTransport& null_transport() { return transport; } -struct Server::Impl { +struct Server::Impl : std::enable_shared_from_this { + enum class LifecycleState : std::uint8_t { + eUninitialized, + eAwaitingInitialized, + eReady, + }; + + struct ToolRegistration { + TypeErasedHandler handler; + detail::ToolResultMode result_mode; + }; + + struct ResourceTemplateRegistration { + ResourceTemplate metadata; + std::string matcher_source; + std::regex matcher; + TypeErasedHandler handler; + }; + Implementation server_info; ServerCapabilities capabilities; - std::atomic_bool initialized{false}; + std::optional instructions; + std::optional discover_ttl_ms; + std::optional discover_cache_scope; + std::atomic lifecycle{LifecycleState::eUninitialized}; std::atomic_bool shutdown_requested{false}; + // Set by ~Server. The implementation itself outlives the Server whenever work is still + // in flight, so this is what tells that work its owner has gone. + std::atomic_bool server_gone{false}; + mutable std::mutex session_mutex; std::shared_ptr session; std::atomic next_request_id{1}; std::atomic log_level{LoggingLevel::eDebug}; @@ -76,12 +389,12 @@ struct Server::Impl { CompletionHandler completion_handler; std::vector tools; - std::map> tool_handlers; + std::map> tool_handlers; std::vector resources; std::map> resource_handlers; - std::vector resource_templates; + std::vector resource_templates; std::vector prompts; std::map> prompt_handlers; @@ -90,35 +403,124 @@ struct Server::Impl { SubscriptionHandler subscribe_handler; SubscriptionHandler unsubscribe_handler; + + // The request path lives here rather than on Server so that work still in flight when a + // Server is destroyed has something valid to run against. Holders keep the implementation + // alive; Server is only the handle the application owns. + [[nodiscard]] std::shared_ptr session_snapshot() const; + void reset_session(const std::shared_ptr& session); + static void abandon_session_work(const std::shared_ptr& session); + + static Task run(std::shared_ptr impl, std::shared_ptr transport, + boost::asio::any_io_executor executor); + Task run_session(std::shared_ptr session); + Task dispatch(nlohmann::json json_msg); + Task notify_resource_updated(const std::string& uri); + Task dispatch_on_strand(nlohmann::json json_msg); + Task dispatch_request(nlohmann::json json_msg); + Task dispatch_request_wire(nlohmann::json json_msg, bool enforce_lifecycle); + void dispatch_notification(const nlohmann::json& json_msg); + void dispatch_response(const nlohmann::json& json_msg); + + Task handle_initialize_wire(const nlohmann::json& json_msg, bool update_lifecycle); + Task handle_shutdown_wire(const nlohmann::json& json_msg); + Task handle_ping_wire(const nlohmann::json& json_msg); + Task handle_discover_wire(const nlohmann::json& json_msg); + Task handle_tools_call_wire(const nlohmann::json& json_msg); + Task handle_tools_list_wire(const nlohmann::json& json_msg); + Task handle_resources_list_wire(const nlohmann::json& json_msg); + Task handle_resources_read_wire(const nlohmann::json& json_msg); + Task handle_resource_templates_list_wire(const nlohmann::json& json_msg); + Task handle_subscribe_wire(const nlohmann::json& json_msg); + Task handle_unsubscribe_wire(const nlohmann::json& json_msg); + Task handle_prompts_list_wire(const nlohmann::json& json_msg); + Task handle_prompts_get_wire(const nlohmann::json& json_msg); + Task handle_set_level_wire(const nlohmann::json& json_msg); + Task handle_complete_wire(const nlohmann::json& json_msg); + + Task invoke_tool_impl(CallToolParams params, + std::shared_ptr> cancelled, + std::optional progress_token); + Context make_context(std::shared_ptr> cancelled = nullptr, + std::optional progress_token = std::nullopt); + TypeErasedHandler build_middleware_chain(TypeErasedHandler final_handler); + + Task send_request(const std::string& method, + const std::optional& params); + static Task await_reverse_response(std::shared_ptr session, + std::shared_ptr wire, + std::int64_t id); + static Task await_reverse_response_on_strand( + std::shared_ptr session, std::shared_ptr wire, std::int64_t id); + Task send_notification(const std::string& method, + const std::optional& params); + Task notify_resource_updated_on_strand(std::shared_ptr session, + std::shared_ptr uri); + + // [gcc11-sso: wire-builders] DO NOT convert to Task. + static std::string make_result_wire(const RequestId& id, nlohmann::json result); + static std::string make_error_wire(const RequestId& id, int code, std::string message); + std::optional paginate(std::size_t total, const nlohmann::json& json_msg); + [[nodiscard]] bool has_tool_output_schema(const std::string& name) const; }; Server::Server(const Implementation& server_info, const ServerCapabilities& capabilities) - : impl_(std::make_unique()) { + : impl_(std::make_shared()) { impl_->server_info = server_info; impl_->capabilities = capabilities; } -Server::~Server() { reset_session(); } +Server::~Server() { + if (impl_) { + impl_->server_gone.store(true, std::memory_order_release); + } + reset_session(); +} -void Server::register_tool(const Tool& tool, const std::string& name, TypeErasedHandler handler) { +void Server::register_tool(const Tool& tool, const std::string& name, TypeErasedHandler handler, + detail::ToolResultMode result_mode) { + if (impl_->tool_handlers.contains(name)) { + throw std::invalid_argument("Duplicate tool registration: " + name); + } impl_->tools.push_back(tool); - impl_->tool_handlers.emplace(name, std::move(handler)); + impl_->tool_handlers.emplace(name, Impl::ToolRegistration{std::move(handler), result_mode}); } void Server::register_resource(const Resource& resource, TypeErasedHandler handler) { auto uri = resource.uri; + if (impl_->resource_handlers.contains(uri)) { + throw std::invalid_argument("Duplicate resource registration: " + uri); + } impl_->resources.push_back(resource); impl_->resource_handlers.emplace(std::move(uri), std::move(handler)); } void Server::register_prompt(const Prompt& prompt, TypeErasedHandler handler) { auto name = prompt.name; + if (impl_->prompt_handlers.contains(name)) { + throw std::invalid_argument("Duplicate prompt registration: " + name); + } impl_->prompts.push_back(prompt); impl_->prompt_handlers.emplace(std::move(name), std::move(handler)); } void Server::add_resource_template(const ResourceTemplate& tmpl) { - impl_->resource_templates.push_back(tmpl); + register_resource_template(tmpl, {}); +} + +void Server::register_resource_template(const ResourceTemplate& tmpl, TypeErasedHandler handler) { + auto pattern = compile_uri_template(tmpl.uriTemplate); + for (const auto& registered : impl_->resource_templates) { + if (registered.metadata.uriTemplate == tmpl.uriTemplate) { + throw std::invalid_argument("Duplicate resource URI template: " + tmpl.uriTemplate); + } + if (registered.matcher_source == pattern.source) { + throw std::invalid_argument("Ambiguous resource URI template: " + tmpl.uriTemplate); + } + } + + impl_->resource_templates.push_back(Impl::ResourceTemplateRegistration{ + tmpl, std::move(pattern.source), std::move(pattern.matcher), std::move(handler)}); } void Server::use(Middleware mw) { impl_->middlewares.push_back(std::move(mw)); } @@ -129,13 +531,28 @@ void Server::set_completion_provider(CompletionHandler handler) { void Server::set_page_size(std::size_t size) { impl_->page_size = size; } +void Server::set_instructions(std::string instructions) { + impl_->instructions = std::move(instructions); +} + +void Server::set_discover_ttl_ms(std::int64_t ttl_ms) { + if (ttl_ms < 0) { + throw std::invalid_argument("Discover ttlMs must be >= 0"); + } + impl_->discover_ttl_ms = ttl_ms; +} + +void Server::set_discover_cache_scope(CacheScope scope) { impl_->discover_cache_scope = scope; } + void Server::on_subscribe(const SubscriptionHandler& handler) { impl_->subscribe_handler = handler; } void Server::on_unsubscribe(const SubscriptionHandler& handler) { impl_->unsubscribe_handler = handler; } -bool Server::is_initialized() const { return impl_->initialized.load(std::memory_order_relaxed); } +bool Server::is_initialized() const { + return impl_->lifecycle.load(std::memory_order_relaxed) != Impl::LifecycleState::eUninitialized; +} bool Server::is_shutdown_requested() const { return impl_->shutdown_requested.load(std::memory_order_relaxed); @@ -143,148 +560,352 @@ bool Server::is_shutdown_requested() const { LoggingLevel Server::get_log_level() const { return impl_->log_level.load(std::memory_order_relaxed); } +// The inverse of the notify_* forwarders below, and safe for the opposite reason: this one returns +// the task instead of awaiting it, which is only sound because Impl::send_request is not a +// coroutine. Its body runs to completion inside this call, so neither reference parameter has to +// outlive the return. Converting Impl::send_request to co_return would leave both dangling here +// with no diagnostic; change this line to co_await in the same commit if you ever do. Task Server::send_request(const std::string& method, const std::optional& params) { - if (!impl_->session || !impl_->session->transport || !impl_->session->strand) { + return impl_->send_request(method, params); +} + +// Deliberately not a coroutine -- see the note on Server::send_request above. It takes both +// parameters by reference and its caller forwards with a plain `return`, so the body must run +// before that return completes. +Task Server::Impl::send_request(const std::string& method, + const std::optional& params) { + // Distinct from the two failures below: a handler still running after its Server was + // destroyed has to be able to tell that apart from a session that is merely closing and + // from stateless dispatch, because the right response differs in each case. + if (server_gone.load(std::memory_order_acquire)) { + throw std::runtime_error("reverse RPC is unavailable: the Server has been destroyed"); + } + auto session = session_snapshot(); + if (!session || !session->transport || !session->strand || + session->stopping.load(std::memory_order_acquire)) { throw std::runtime_error("reverse RPC is unavailable in stateless direct dispatch"); } // [gcc11-sso: int64-id] DO NOT change id to std::string. - int64_t id = impl_->next_request_id.fetch_add(1, std::memory_order_relaxed); - boost::asio::steady_timer* timer_ptr = nullptr; - std::string wire; - { - auto id_str = std::to_string(id); - JSONRPCRequest request; - request.id = RequestId{id_str}; - request.method = method; - request.params = params; + const int64_t id = next_request_id.fetch_add(1, std::memory_order_relaxed); + JSONRPCRequest request; + request.id = RequestId{std::to_string(id)}; + request.method = method; + request.params = params; + auto wire = std::make_shared(nlohmann::json(std::move(request)).dump()); + return await_reverse_response(std::move(session), std::move(wire), id); +} - auto& pending = impl_->session->pending_requests[id_str]; - pending.timer = std::make_unique(*impl_->session->strand); - pending.timer->expires_at(std::chrono::steady_clock::time_point::max()); - timer_ptr = pending.timer.get(); +Task Server::Impl::await_reverse_response(std::shared_ptr session, + std::shared_ptr wire, + int64_t id) { + auto strand = *session->strand; + // A caller's awaitable keeps its original executor across an awaited post. Launch the + // complete correlation lifecycle on the session strand so map and timer state stay confined. + return boost::asio::co_spawn( + strand, await_reverse_response_on_strand(std::move(session), std::move(wire), id), + boost::asio::use_awaitable); +} - wire = nlohmann::json(request).dump(); - } // id_str, request, method destroyed here — before co_await +Task Server::Impl::await_reverse_response_on_strand( + std::shared_ptr session, std::shared_ptr wire, int64_t id) { + if (session->stopping.load(std::memory_order_acquire)) { + throw std::runtime_error("server session is closing"); + } - co_await impl_->session->transport->write_message(wire); + auto timer = std::make_shared(*session->strand); + timer->expires_at(std::chrono::steady_clock::time_point::max()); + { + const auto id_key = RequestId{std::to_string(id)}.correlation_key(); + const auto [pending_it, inserted] = session->pending_requests.try_emplace( + id_key, PendingRequest{timer, {}, std::nullopt, false}); + static_cast(pending_it); + if (!inserted) { + throw std::runtime_error("duplicate pending request id: " + id_key); + } + } try { - co_await timer_ptr->async_wait(boost::asio::use_awaitable); - } catch (const boost::system::system_error& err) { - if (err.code() != boost::asio::error::operation_aborted) { - throw; + co_await session->writer->write_message(wire); + } catch (...) { + session->pending_requests.erase(RequestId{std::to_string(id)}.correlation_key()); + throw; + } + + bool completed = false; + { + const auto id_key = RequestId{std::to_string(id)}.correlation_key(); + const auto pending_it = session->pending_requests.find(id_key); + completed = pending_it != session->pending_requests.end() && pending_it->second.completed; + } + if (!completed) { + try { + co_await timer->async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& err) { + if (err.code() != boost::asio::error::operation_aborted) { + session->pending_requests.erase(RequestId{std::to_string(id)}.correlation_key()); + throw; + } } } - // [gcc11-sso: int64-id] Rebuild id_str after all suspensions. - auto id_str = std::to_string(id); - auto it = impl_->session->pending_requests.find(id_str); - if (it == impl_->session->pending_requests.end()) { - throw std::runtime_error("pending request not found for id: " + id_str); + // [gcc11-sso: int64-id] Rebuild the string only after all suspensions. + const auto id_key = RequestId{std::to_string(id)}.correlation_key(); + auto it = session->pending_requests.find(id_key); + if (it == session->pending_requests.end()) { + throw std::runtime_error("pending request not found for id: " + std::to_string(id)); } auto json_result = std::move(it->second.result); auto error = std::move(it->second.error); - impl_->session->pending_requests.erase(it); + session->pending_requests.erase(it); if (error) { + // Same defect as the peer-controlled text in the client's McpError, opposite direction: + // this is the error a CLIENT returned for a server-initiated request (sampling, + // elicitation, roots/list), stored verbatim by dispatch_response(). `message` is the + // client's text and this diagnostic is what the server operator logs, with no + // authorization step in the way, so it is flattened and bounded before it goes in. throw std::runtime_error("JSON-RPC error " + std::to_string(error->code) + ": " + - error->message); + detail::sanitize_for_diagnostics(error->message)); } co_return json_result; } Task Server::invoke_tool(const std::string& tool_name, const nlohmann::json& args) { - co_return co_await invoke_tool_impl(CallToolParams{tool_name, args, std::nullopt}, nullptr, - std::nullopt); + return impl_->invoke_tool_impl(CallToolParams{tool_name, args, std::nullopt}, nullptr, + std::nullopt); } Task Server::run(std::shared_ptr transport, boost::asio::any_io_executor executor) { + // Deliberately not a coroutine: reading impl_ here, on the caller's thread, hands the session + // a reference that outlives this Server rather than one that is read again later. + return Impl::run(impl_, std::move(transport), std::move(executor)); +} + +Task Server::Impl::run(std::shared_ptr impl, std::shared_ptr transport, + boost::asio::any_io_executor executor) { + if (!transport) { + throw std::invalid_argument("Server transport must not be null"); + } + auto session = std::make_shared(); session->transport = std::move(transport); session->memory_transport = dynamic_cast(session->transport.get()); session->strand = std::make_unique>( boost::asio::make_strand(executor)); - impl_->session = session; + session->writer = + std::make_shared(session->transport, *session->strand); + session->drain_timer = std::make_shared(*session->strand); + session->drain_timer->expires_at(std::chrono::steady_clock::time_point::max()); + auto strand = *session->strand; + // The read loop and teardown both own session state, so run both on the session strand. + co_await boost::asio::co_spawn(strand, impl->run_session(std::move(session)), + boost::asio::use_awaitable); +} + +Task Server::Impl::run_session(std::shared_ptr new_session) { + auto session = std::move(new_session); + { + std::lock_guard lock(session_mutex); + if (this->session) { + throw std::runtime_error("Server already has an active session"); + } + this->session = session; + } + + lifecycle.store(Impl::LifecycleState::eUninitialized, std::memory_order_relaxed); + shutdown_requested.store(false, std::memory_order_relaxed); try { for (;;) { nlohmann::json json_msg; - if (session->memory_transport) { - json_msg = co_await session->memory_transport->read_json(); - } else { - auto raw = co_await session->transport->read_message(); - json_msg = nlohmann::json::parse(raw); + bool parse_failed = false; + try { + if (session->memory_transport) { + json_msg = co_await session->memory_transport->read_json(); + } else { + auto raw = co_await session->transport->read_message(); + json_msg = nlohmann::json::parse(raw); + } + } catch (const nlohmann::json::parse_error&) { + parse_failed = true; + } + if (parse_failed) { + co_await session->writer->write_message(make_parse_error_wire()); + continue; } + session->active_dispatches.fetch_add(1, std::memory_order_relaxed); boost::asio::co_spawn( *session->strand, - [this, session, json_msg = std::move(json_msg)]() mutable -> Task { - co_await dispatch(std::move(json_msg)); + [impl = shared_from_this(), session, + json_msg = std::move(json_msg)]() mutable -> Task { + try { + co_await impl->dispatch_on_strand(std::move(json_msg)); + } catch (...) { + // A failed request must not terminate the detached dispatcher. + } + + const auto remaining = + session->active_dispatches.fetch_sub(1, std::memory_order_acq_rel) - 1; + if (remaining == 0 && session->stopping.load(std::memory_order_acquire)) { + session->drain_timer->expires_at(std::chrono::steady_clock::now()); + } co_return; }, boost::asio::detached); } - } catch (const std::exception& e) { - (void)e; + } catch (const std::exception&) { + // Closing a transport is the normal way to stop a session. } - reset_session(); + // Teardown below can fail — a transport may throw from close, and the drain wait rethrows any + // error that is not cancellation. Whichever way it ends, the session must be unregistered, or + // every later run() is refused for the lifetime of this Server. + try { + session->stopping.store(true, std::memory_order_release); + abandon_session_work(session); + session->transport->close(); + + while (session->active_dispatches.load(std::memory_order_acquire) != 0) { + session->drain_timer->expires_at(std::chrono::steady_clock::time_point::max()); + if (session->active_dispatches.load(std::memory_order_acquire) == 0) { + break; + } + try { + co_await session->drain_timer->async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + } + } catch (...) { + reset_session(session); + throw; + } + + reset_session(session); } -Task Server::dispatch(nlohmann::json json_msg) { - if (!json_msg.contains("id")) { +Task Server::dispatch(nlohmann::json json_msg) { return impl_->dispatch(std::move(json_msg)); } + +Task Server::Impl::dispatch(nlohmann::json json_msg) { + auto session = session_snapshot(); + if (session && session->strand) { + auto strand = *session->strand; + co_await boost::asio::co_spawn(strand, dispatch_on_strand(std::move(json_msg)), + boost::asio::use_awaitable); + co_return; + } + + co_await dispatch_on_strand(std::move(json_msg)); +} + +Task Server::Impl::dispatch_on_strand(nlohmann::json json_msg) { + if (is_valid_notification_envelope(json_msg)) { dispatch_notification(json_msg); co_return; } - if (!json_msg.contains("method")) { + // A method without an id is notification-shaped. Notifications never receive responses, + // including when their envelope is malformed. + if (json_msg.is_object() && json_msg.contains("method") && !json_msg.contains("id")) { + co_return; + } + + if (json_msg.is_object() && !json_msg.contains("method") && + (json_msg.contains("result") || has_error_member(json_msg))) { dispatch_response(json_msg); co_return; } + // Anything that is neither a notification nor a response is request-shaped. Route invalid + // values through request validation so the peer receives -32600 instead of a silent drop. co_await dispatch_request(std::move(json_msg)); } Task Server::dispatch_request_direct(nlohmann::json json_msg) { - co_return co_await dispatch_request_wire(std::move(json_msg)); + return impl_->dispatch_request_wire(std::move(json_msg), false); } -Context Server::make_context(std::shared_ptr> cancelled, - std::optional progress_token) { +Context Server::Impl::make_context(std::shared_ptr> cancelled, + std::optional progress_token) { ITransport* transport = &null_transport(); - if (impl_->session && impl_->session->transport) { - transport = impl_->session->transport.get(); + MessageSender message_sender; + auto session = session_snapshot(); + if (session && session->transport) { + transport = session->transport.get(); + auto writer = session->writer; + message_sender = [writer = std::move(writer)](std::string_view message) { + return writer->write_message(message); + }; } return {*transport, - [this](std::string method, std::optional params) -> Task { - co_return co_await send_request(std::move(method), std::move(params)); + [impl = shared_from_this()](std::string method, + std::optional params) -> Task { + co_return co_await impl->send_request(std::move(method), std::move(params)); }, - std::move(cancelled), std::move(progress_token), &impl_->log_level}; + std::move(cancelled), + std::move(progress_token), + &log_level, + std::move(message_sender)}; } -// Exceptions from handlers are caught and reported as g_INTERNAL_ERROR (-32603) responses. -Task Server::dispatch_request(nlohmann::json json_msg) { - co_await impl_->session->transport->write_message( - co_await dispatch_request_wire(std::move(json_msg))); +// Invalid parameter decoding is reported as -32602. Exceptions raised after decoding are reported +// as -32603, except from a tool handler: those become a tool result carrying isError. +Task Server::Impl::dispatch_request(nlohmann::json json_msg) { + auto session = session_snapshot(); + if (!session || !session->writer) { + throw std::runtime_error("server dispatch requires an active session"); + } + co_await session->writer->write_message(co_await dispatch_request_wire(std::move(json_msg), true)); } -Task Server::dispatch_request_wire(nlohmann::json json_msg) { +Task Server::Impl::dispatch_request_wire(nlohmann::json json_msg, bool enforce_lifecycle) { + { + const char* validation_error = validate_request_envelope(json_msg); + if (validation_error != nullptr) { + co_return make_invalid_request_wire(json_msg, validation_error); + } + } + // [gcc11-sso: string_view] DO NOT change to std::string. std::string_view method = json_msg.at("method").get_ref(); // [gcc11-sso: scope-before-await] DO NOT use optional for error state. nlohmann::json error_payload; + int error_code = g_INTERNAL_ERROR; try { if (method == "initialize") { - co_return co_await handle_initialize_wire(json_msg); - } else if (method == "ping") { + if (enforce_lifecycle && + lifecycle.load(std::memory_order_relaxed) != Impl::LifecycleState::eUninitialized) { + co_return make_error_wire(json_msg.at("id").get(), g_INVALID_REQUEST, + "Server has already been initialized"); + } + co_return co_await handle_initialize_wire(json_msg, enforce_lifecycle); + } + + if (method == "ping") { co_return co_await handle_ping_wire(json_msg); - } else if (method == "shutdown") { + } + + if (method == "server/discover") { + co_return co_await handle_discover_wire(json_msg); + } + + if (enforce_lifecycle && + lifecycle.load(std::memory_order_relaxed) != Impl::LifecycleState::eReady) { + co_return make_error_wire( + json_msg.at("id").get(), g_INVALID_REQUEST, + "Server is not ready; initialize and send notifications/initialized first"); + } + + if (method == "shutdown") { co_return co_await handle_shutdown_wire(json_msg); } else if (method == "tools/call") { co_return co_await handle_tools_call_wire(json_msg); @@ -312,14 +933,19 @@ Task Server::dispatch_request_wire(nlohmann::json json_msg) { co_return make_error_wire(json_msg.at("id").get(), g_METHOD_NOT_FOUND, "Method not found: " + std::string(method)); } - } catch (const std::exception& e) { + } catch (const InvalidParamsError& error) { + error_code = g_INVALID_PARAMS; + error_payload = error.what(); + } catch (const std::exception& error) { if (json_msg.contains("id")) { - error_payload = e.what(); + // Handler text can carry build paths or third-party library detail, and reaches the + // peer verbatim from here. + error_payload = detail::sanitize_for_diagnostics(error.what()); } } if (!error_payload.is_null()) { - co_return make_error_wire(json_msg.at("id").get(), g_INTERNAL_ERROR, + co_return make_error_wire(json_msg.at("id").get(), error_code, error_payload.get()); } @@ -327,167 +953,233 @@ Task Server::dispatch_request_wire(nlohmann::json json_msg) { "Request produced no response"); } -void Server::dispatch_notification(const nlohmann::json& json_msg) { - if (!json_msg.contains("method")) { +void Server::Impl::dispatch_notification(const nlohmann::json& json_msg) { + if (!is_valid_notification_envelope(json_msg)) { return; } auto method = json_msg.at("method").get(); + if (method == "notifications/initialized") { + auto expected = Impl::LifecycleState::eAwaitingInitialized; + lifecycle.compare_exchange_strong(expected, Impl::LifecycleState::eReady, + std::memory_order_relaxed); + return; + } + if (method == "notifications/cancelled") { + auto session = session_snapshot(); + if (lifecycle.load(std::memory_order_relaxed) != Impl::LifecycleState::eReady || !session) { + return; + } if (!json_msg.contains("params")) { return; } - auto params = json_msg.at("params").get(); - auto id_str = params.requestId.to_string(); + CancelledNotificationParams params; + try { + params = json_msg.at("params").get(); + } catch (const std::exception&) { + return; + } + auto id_str = params.requestId.correlation_key(); - auto it = impl_->session->in_flight.find(id_str); - if (it != impl_->session->in_flight.end()) { + auto it = session->in_flight.find(id_str); + if (it != session->in_flight.end()) { it->second->store(true, std::memory_order_relaxed); } } } -void Server::dispatch_response(const nlohmann::json& json_msg) { +void Server::Impl::dispatch_response(const nlohmann::json& json_msg) { + auto session = session_snapshot(); + if (!is_valid_response_envelope(json_msg) || !session) { + return; + } + auto id = json_msg.at("id").get(); - auto id_str = id.to_string(); + auto id_str = id.correlation_key(); - auto it = impl_->session->pending_requests.find(id_str); - if (it == impl_->session->pending_requests.end()) { + auto it = session->pending_requests.find(id_str); + if (it == session->pending_requests.end()) { return; } - if (json_msg.contains("error")) { + if (has_error_member(json_msg)) { it->second.error = json_msg.at("error").get(); } else if (json_msg.contains("result")) { it->second.result = json_msg.at("result"); } + it->second.completed = true; it->second.timer->cancel(); } -Task Server::handle_initialize(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_initialize_wire(json_msg)); -} +Task Server::Impl::handle_initialize_wire(const nlohmann::json& json_msg, + bool update_lifecycle) { + auto initialize_request = deserialize_request_params(json_msg, "initialize"); -Task Server::handle_initialize_wire(const nlohmann::json& json_msg) { InitializeResult init_result; - std::string negotiated_protocol_version = std::string(g_LATEST_PROTOCOL_VERSION); - if (json_msg.contains("params") && json_msg.at("params").is_object()) { - const auto& params = json_msg.at("params"); - if (params.contains("protocolVersion") && params.at("protocolVersion").is_string()) { - const auto requested_protocol_version = params.at("protocolVersion").get(); - negotiated_protocol_version = - std::string(negotiate_protocol_version(requested_protocol_version)); - } + init_result.protocolVersion = + std::string(negotiate_protocol_version(initialize_request.protocolVersion)); + init_result.capabilities = capabilities; + init_result.serverInfo = server_info; + init_result.instructions = instructions; + + if (update_lifecycle) { + lifecycle.store(Impl::LifecycleState::eAwaitingInitialized, std::memory_order_relaxed); } - - init_result.protocolVersion = std::move(negotiated_protocol_version); - init_result.capabilities = impl_->capabilities; - init_result.serverInfo = impl_->server_info; - - impl_->initialized.store(true, std::memory_order_relaxed); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json(std::move(init_result))); } -Task Server::handle_shutdown(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_shutdown_wire(json_msg)); -} - -Task Server::handle_shutdown_wire(const nlohmann::json& json_msg) { - impl_->shutdown_requested.store(true, std::memory_order_relaxed); +Task Server::Impl::handle_shutdown_wire(const nlohmann::json& json_msg) { + shutdown_requested.store(true, std::memory_order_relaxed); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json::object()); } -Task Server::handle_ping(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_ping_wire(json_msg)); +Task Server::Impl::handle_ping_wire(const nlohmann::json& json_msg) { + co_return make_result_wire(json_msg.at("id").get(), nlohmann::json::object()); } -Task Server::handle_ping_wire(const nlohmann::json& json_msg) { - co_return make_result_wire(json_msg.at("id").get(), nlohmann::json::object()); +// Params carry only `_meta`, whose contents (protocolVersion, clientInfo, clientCapabilities) +// are accepted but not yet interpreted. This request is a pre-gate method +// like initialize/ping: reachable with no prior state and idempotent, so it neither reads +// json_msg's params nor mutates lifecycle. +Task Server::Impl::handle_discover_wire(const nlohmann::json& json_msg) { + DiscoverResult discover_result; + // resultType uses the DiscoverResult struct default ("complete"); a future revision introduces + // a shared result-envelope helper for this field that other cacheable results will also use. + discover_result.supportedVersions.assign(g_DISCOVERABLE_PROTOCOL_VERSIONS.begin(), + g_DISCOVERABLE_PROTOCOL_VERSIONS.end()); + discover_result.capabilities = capabilities; + discover_result.serverInfo = server_info; + discover_result.instructions = instructions; + // server/utilities/caching.md: servers MUST include caching hints on "complete" results, + // server/discover listed first, and MUST provide a ttlMs >= 0. ttlMs defaults to 0 ("do not + // cache" / immediately stale per the spec's freshness rule); an absent ttlMs "should only + // occur in older server versions", which this SDK is not, so it is always emitted. The spec + // does not state a default cacheScope, so this SDK defaults to the conservative choice, + // "private" (do not assume the result is safe to share across authorization contexts), + // until a caller explicitly opts into "public" via set_discover_cache_scope. + discover_result.ttlMs = discover_ttl_ms.value_or(0); + discover_result.cacheScope = discover_cache_scope.value_or(CacheScope::ePrivate); + co_return make_result_wire(json_msg.at("id").get(), nlohmann::json(discover_result)); } -Task Server::invoke_tool_impl(CallToolParams params, - std::shared_ptr> cancelled, - std::optional progress_token) { - auto iter = impl_->tool_handlers.find(params.name); - if (iter == impl_->tool_handlers.end()) { - throw std::runtime_error("Unknown tool: " + params.name); +Task Server::Impl::invoke_tool_impl(CallToolParams params, + std::shared_ptr> cancelled, + std::optional progress_token) { + auto iter = tool_handlers.find(params.name); + if (iter == tool_handlers.end()) { + // The name is chosen by whoever called: over JSON-RPC that is the client, and through the + // public `invoke_tool` entry point it is whatever text the embedding passed in. Flatten and + // bound it, or a name carrying CR/LF forges a line in the operator's log. + throw std::runtime_error("Unknown tool: " + detail::sanitize_for_diagnostics(params.name)); } auto ctx = make_context(std::move(cancelled), std::move(progress_token)); nlohmann::json handler_result; - if (impl_->middlewares.empty()) { - handler_result = co_await iter->second(ctx, params.arguments); + // A tool that throws has failed at its own job, which the protocol reports inside the result as + // isError. Middleware is not the tool: it decides whether the call may proceed at all, so it + // runs outside this catch and a middleware failure surfaces as a JSON-RPC error instead. + TypeErasedHandler guarded_handler = [original_handler = iter->second.handler]( + Context& inner_ctx, + const nlohmann::json& arguments) -> Task { + try { + co_return co_await original_handler(inner_ctx, arguments); + } catch (const std::exception& error) { + co_return nlohmann::json( + make_tool_error_result(detail::sanitize_for_diagnostics(error.what()))); + } catch (...) { + // Nothing above catches a throw that does not derive from std::exception, and nothing + // further out does either: handle_tools_call_wire rethrows and dispatch_request_wire + // has no catch-all, so the throw escapes the dispatcher, no response is ever written + // and the caller waits out its own request timeout. The text is ours rather than the + // thrower's, so there is nothing to sanitize. + co_return nlohmann::json(make_tool_error_result("Tool handler failed")); + } + }; + + if (middlewares.empty()) { + handler_result = co_await guarded_handler(ctx, params.arguments); } else { TypeErasedHandler wrapped_handler = - [original_handler = iter->second]( - Context& inner_ctx, const nlohmann::json& full_params) -> Task { + [guarded_handler](Context& inner_ctx, + const nlohmann::json& full_params) -> Task { auto call_params = full_params.get(); - co_return co_await original_handler(inner_ctx, call_params.arguments); + co_return co_await guarded_handler(inner_ctx, call_params.arguments); }; auto handler = build_middleware_chain(std::move(wrapped_handler)); nlohmann::json params_json = params; handler_result = co_await handler(ctx, params_json); } - if (has_tool_output_schema(params.name)) { - nlohmann::json structured = handler_result; - handler_result["structuredContent"] = std::move(structured); + if (iter->second.result_mode == detail::ToolResultMode::eValidated) { + validate_call_tool_result(handler_result); + } else { + handler_result = normalize_structured_tool_result(std::move(handler_result)); } - co_return handler_result; -} + const bool is_error_result = handler_result.value("isError", false); + if (has_tool_output_schema(params.name) && !is_error_result && + !handler_result.contains("structuredContent")) { + throw std::invalid_argument("Tool with outputSchema must return structuredContent: " + + params.name); + } -Task Server::handle_tools_call(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_tools_call_wire(json_msg)); + co_return handler_result; } -Task Server::handle_tools_call_wire(const nlohmann::json& json_msg) { - auto params = json_msg.at("params").get(); - if (!impl_->tool_handlers.contains(params.name)) { +Task Server::Impl::handle_tools_call_wire(const nlohmann::json& json_msg) { + auto params = deserialize_request_params(json_msg, "tools/call"); + if (!tool_handlers.contains(params.name)) { + // This, not the throw in invoke_tool_impl, is the site a remote client actually reaches: + // the RPC path refuses before dispatch. The name is entirely the client's, so flatten and + // bound it before it reaches the operator's log. co_return make_error_wire(json_msg.at("id").get(), g_METHOD_NOT_FOUND, - "Unknown tool: " + params.name); + "Unknown tool: " + detail::sanitize_for_diagnostics(params.name)); } // [gcc11-sso: int64-id] request_id_str from json_msg (heap-safe); used only in in_flight map. - auto request_id_str = json_msg.at("id").get().to_string(); - - auto cancelled = std::make_shared>(false); - if (impl_->session) { - impl_->session->in_flight[request_id_str] = cancelled; - } + auto request_id_str = json_msg.at("id").get().correlation_key(); std::optional progress_token; - if (params.meta) { - auto& meta = *params.meta; - if (meta.contains("progressToken")) { - progress_token = meta.at("progressToken").get(); + try { + if (params.meta) { + auto& meta = *params.meta; + if (meta.contains("progressToken")) { + progress_token = meta.at("progressToken").get(); + } } + } catch (const nlohmann::json::exception& error) { + throw InvalidParamsError("Invalid tools/call params: " + std::string(error.what())); + } catch (const std::invalid_argument& error) { + throw InvalidParamsError("Invalid tools/call params: " + std::string(error.what())); + } + + auto cancelled = std::make_shared>(false); + auto session = session_snapshot(); + if (session) { + session->in_flight[request_id_str] = cancelled; } try { nlohmann::json handler_result = co_await invoke_tool_impl(std::move(params), cancelled, std::move(progress_token)); - if (impl_->session) { - impl_->session->in_flight.erase(request_id_str); + if (session) { + session->in_flight.erase(request_id_str); } co_return make_result_wire(json_msg.at("id").get(), std::move(handler_result)); } catch (...) { - if (impl_->session) { - impl_->session->in_flight.erase(request_id_str); + if (session) { + session->in_flight.erase(request_id_str); } throw; } } -Task Server::handle_tools_list(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_tools_list_wire(json_msg)); -} - -Task Server::handle_tools_list_wire(const nlohmann::json& json_msg) { - auto page = paginate(impl_->tools.size(), json_msg); +Task Server::Impl::handle_tools_list_wire(const nlohmann::json& json_msg) { + auto page = paginate(tools.size(), json_msg); if (!page) { co_return make_error_wire(json_msg.at("id").get(), g_INVALID_PARAMS, "Invalid pagination cursor"); @@ -495,20 +1187,16 @@ Task Server::handle_tools_list_wire(const nlohmann::json& json_msg) ListToolsResult list_result; auto [begin, end, next_cursor] = *page; - list_result.tools.assign(impl_->tools.begin() + static_cast(begin), - impl_->tools.begin() + static_cast(end)); + list_result.tools.assign(tools.begin() + static_cast(begin), + tools.begin() + static_cast(end)); list_result.nextCursor = std::move(next_cursor); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json(std::move(list_result))); } -Task Server::handle_resources_list(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_resources_list_wire(json_msg)); -} - -Task Server::handle_resources_list_wire(const nlohmann::json& json_msg) { - auto page = paginate(impl_->resources.size(), json_msg); +Task Server::Impl::handle_resources_list_wire(const nlohmann::json& json_msg) { + auto page = paginate(resources.size(), json_msg); if (!page) { co_return make_error_wire(json_msg.at("id").get(), g_INVALID_PARAMS, "Invalid pagination cursor"); @@ -516,41 +1204,66 @@ Task Server::handle_resources_list_wire(const nlohmann::json& json_ ListResourcesResult list_result; auto [begin, end, next_cursor] = *page; - list_result.resources.assign(impl_->resources.begin() + static_cast(begin), - impl_->resources.begin() + static_cast(end)); + list_result.resources.assign(resources.begin() + static_cast(begin), + resources.begin() + static_cast(end)); list_result.nextCursor = std::move(next_cursor); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json(std::move(list_result))); } -Task Server::handle_resources_read(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_resources_read_wire(json_msg)); -} +Task Server::Impl::handle_resources_read_wire(const nlohmann::json& json_msg) { + auto params = deserialize_request_params(json_msg, "resources/read"); + auto iter = resource_handlers.find(params.uri); + if (iter != resource_handlers.end()) { + auto handler = build_middleware_chain(iter->second); + + nlohmann::json params_json = params; + auto ctx = make_context(); + nlohmann::json handler_result = co_await handler(ctx, params_json); + co_return make_result_wire(json_msg.at("id").get(), std::move(handler_result)); + } + + // Exact resources are found by lookup and so may be any length; only template matching is + // length-sensitive. This reports the limit rather than "Unknown resource", so that a caller + // whose URI is merely too long can tell that from one that names nothing. + if (params.uri.size() > g_MAX_TEMPLATE_MATCH_URI_LENGTH) { + co_return make_error_wire(json_msg.at("id").get(), g_INVALID_PARAMS, + "Resource URI exceeds the " + + std::to_string(g_MAX_TEMPLATE_MATCH_URI_LENGTH) + + " character limit for template matching"); + } + + const Impl::ResourceTemplateRegistration* match = nullptr; + for (const auto& resource_template : resource_templates) { + if (!resource_template.handler) { + continue; + } + if (!std::regex_match(params.uri, resource_template.matcher)) { + continue; + } + if (match != nullptr) { + co_return make_error_wire(json_msg.at("id").get(), g_INVALID_PARAMS, + "Ambiguous resource template match: " + params.uri); + } + match = &resource_template; + } -Task Server::handle_resources_read_wire(const nlohmann::json& json_msg) { - auto params = json_msg.at("params").get(); - auto iter = impl_->resource_handlers.find(params.uri); - if (iter == impl_->resource_handlers.end()) { + if (match == nullptr || !match->handler) { co_return make_error_wire(json_msg.at("id").get(), g_INVALID_PARAMS, "Unknown resource: " + params.uri); } - auto handler = build_middleware_chain(iter->second); + auto handler = build_middleware_chain(match->handler); - nlohmann::json params_json = std::move(params); + nlohmann::json params_json = params; auto ctx = make_context(); nlohmann::json handler_result = co_await handler(ctx, params_json); co_return make_result_wire(json_msg.at("id").get(), std::move(handler_result)); } -Task Server::handle_resource_templates_list(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message( - co_await handle_resource_templates_list_wire(json_msg)); -} - -Task Server::handle_resource_templates_list_wire(const nlohmann::json& json_msg) { - auto page = paginate(impl_->resource_templates.size(), json_msg); +Task Server::Impl::handle_resource_templates_list_wire(const nlohmann::json& json_msg) { + auto page = paginate(resource_templates.size(), json_msg); if (!page) { co_return make_error_wire(json_msg.at("id").get(), g_INVALID_PARAMS, "Invalid pagination cursor"); @@ -558,51 +1271,43 @@ Task Server::handle_resource_templates_list_wire(const nlohmann::js ListResourceTemplatesResult list_result; auto [begin, end, next_cursor] = *page; - list_result.resourceTemplates.assign( - impl_->resource_templates.begin() + static_cast(begin), - impl_->resource_templates.begin() + static_cast(end)); + list_result.resourceTemplates.reserve(end - begin); + for (std::size_t index = begin; index < end; ++index) { + list_result.resourceTemplates.push_back(resource_templates[index].metadata); + } list_result.nextCursor = std::move(next_cursor); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json(std::move(list_result))); } -Task Server::handle_subscribe(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_subscribe_wire(json_msg)); -} - -Task Server::handle_subscribe_wire(const nlohmann::json& json_msg) { - auto params = json_msg.at("params").get(); - if (impl_->session) { - impl_->session->subscriptions[params.uri] = true; +Task Server::Impl::handle_subscribe_wire(const nlohmann::json& json_msg) { + auto params = deserialize_request_params(json_msg, "resources/subscribe"); + auto session = session_snapshot(); + if (session) { + session->subscriptions[params.uri] = true; } - if (impl_->subscribe_handler) { - impl_->subscribe_handler(params.uri); + if (subscribe_handler) { + subscribe_handler(params.uri); } co_return make_result_wire(json_msg.at("id").get(), nlohmann::json::object()); } -Task Server::handle_unsubscribe(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_unsubscribe_wire(json_msg)); -} - -Task Server::handle_unsubscribe_wire(const nlohmann::json& json_msg) { - auto params = json_msg.at("params").get(); - if (impl_->session) { - impl_->session->subscriptions.erase(params.uri); +Task Server::Impl::handle_unsubscribe_wire(const nlohmann::json& json_msg) { + auto params = + deserialize_request_params(json_msg, "resources/unsubscribe"); + auto session = session_snapshot(); + if (session) { + session->subscriptions.erase(params.uri); } - if (impl_->unsubscribe_handler) { - impl_->unsubscribe_handler(params.uri); + if (unsubscribe_handler) { + unsubscribe_handler(params.uri); } co_return make_result_wire(json_msg.at("id").get(), nlohmann::json::object()); } -Task Server::handle_prompts_list(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_prompts_list_wire(json_msg)); -} - -Task Server::handle_prompts_list_wire(const nlohmann::json& json_msg) { - auto page = paginate(impl_->prompts.size(), json_msg); +Task Server::Impl::handle_prompts_list_wire(const nlohmann::json& json_msg) { + auto page = paginate(prompts.size(), json_msg); if (!page) { co_return make_error_wire(json_msg.at("id").get(), g_INVALID_PARAMS, "Invalid pagination cursor"); @@ -610,22 +1315,18 @@ Task Server::handle_prompts_list_wire(const nlohmann::json& json_ms ListPromptsResult list_result; auto [begin, end, next_cursor] = *page; - list_result.prompts.assign(impl_->prompts.begin() + static_cast(begin), - impl_->prompts.begin() + static_cast(end)); + list_result.prompts.assign(prompts.begin() + static_cast(begin), + prompts.begin() + static_cast(end)); list_result.nextCursor = std::move(next_cursor); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json(std::move(list_result))); } -Task Server::handle_prompts_get(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_prompts_get_wire(json_msg)); -} - -Task Server::handle_prompts_get_wire(const nlohmann::json& json_msg) { - auto params = json_msg.at("params").get(); - auto iter = impl_->prompt_handlers.find(params.name); - if (iter == impl_->prompt_handlers.end()) { +Task Server::Impl::handle_prompts_get_wire(const nlohmann::json& json_msg) { + auto params = deserialize_request_params(json_msg, "prompts/get"); + auto iter = prompt_handlers.find(params.name); + if (iter == prompt_handlers.end()) { co_return make_error_wire(json_msg.at("id").get(), g_METHOD_NOT_FOUND, "Unknown prompt: " + params.name); } @@ -638,41 +1339,34 @@ Task Server::handle_prompts_get_wire(const nlohmann::json& json_msg co_return make_result_wire(json_msg.at("id").get(), std::move(handler_result)); } -Task Server::handle_set_level(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_set_level_wire(json_msg)); -} - -Task Server::handle_set_level_wire(const nlohmann::json& json_msg) { - auto params = json_msg.at("params").get(); - impl_->log_level.store(params.level, std::memory_order_relaxed); +Task Server::Impl::handle_set_level_wire(const nlohmann::json& json_msg) { + auto params = deserialize_request_params(json_msg, "logging/setLevel"); + log_level.store(params.level, std::memory_order_relaxed); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json::object()); } -Task Server::handle_complete(const nlohmann::json& json_msg) { - co_await impl_->session->transport->write_message(co_await handle_complete_wire(json_msg)); -} +Task Server::Impl::handle_complete_wire(const nlohmann::json& json_msg) { + auto params = deserialize_request_params(json_msg, "completion/complete"); -Task Server::handle_complete_wire(const nlohmann::json& json_msg) { - if (!impl_->completion_handler) { + if (!completion_handler) { co_return make_error_wire(json_msg.at("id").get(), g_METHOD_NOT_FOUND, "No completion handler registered"); } - auto params = json_msg.at("params").get(); - auto complete_result = co_await impl_->completion_handler(params); + auto complete_result = co_await completion_handler(params); co_return make_result_wire(json_msg.at("id").get(), nlohmann::json(std::move(complete_result))); } -std::string Server::make_result_wire(const RequestId& id, nlohmann::json result) { +std::string Server::Impl::make_result_wire(const RequestId& id, nlohmann::json result) { JSONRPCResultResponse response; response.id = id; response.result = std::move(result); return nlohmann::json(std::move(response)).dump(); } -std::string Server::make_error_wire(const RequestId& id, int code, std::string message) { +std::string Server::Impl::make_error_wire(const RequestId& id, int code, std::string message) { Error error; error.code = code; error.message = std::move(message); @@ -682,9 +1376,10 @@ std::string Server::make_error_wire(const RequestId& id, int code, std::string m return nlohmann::json(std::move(response)).dump(); } -Task Server::send_notification(const std::string& method, - const std::optional& params) { - if (!impl_->session) { +Task Server::Impl::send_notification(const std::string& method, + const std::optional& params) { + auto session = session_snapshot(); + if (!session || session->stopping.load(std::memory_order_acquire)) { co_return; } // [gcc11-sso: scope-before-await] @@ -695,32 +1390,55 @@ Task Server::send_notification(const std::string& method, notification.params = params; wire = nlohmann::json(std::move(notification)).dump(); } - co_await impl_->session->transport->write_message(wire); + co_await session->writer->write_message(wire); } +// These three await rather than returning the task, and that is load-bearing. The argument is a +// temporary bound to a reference parameter of a lazy coroutine: under co_await it lives to the end +// of the full expression, which includes the await, but a plain `return` destroys it before the +// coroutine body ever runs and leaves the reference dangling. Both forms compile. Task Server::notify_tools_list_changed() { - co_await send_notification("notifications/tools/list_changed", std::nullopt); + co_await impl_->send_notification("notifications/tools/list_changed", std::nullopt); } Task Server::notify_resources_list_changed() { - co_await send_notification("notifications/resources/list_changed", std::nullopt); + co_await impl_->send_notification("notifications/resources/list_changed", std::nullopt); } Task Server::notify_prompts_list_changed() { - co_await send_notification("notifications/prompts/list_changed", std::nullopt); + co_await impl_->send_notification("notifications/prompts/list_changed", std::nullopt); } Task Server::notify_resource_updated(const std::string& uri) { - if (!impl_->session) { + co_await impl_->notify_resource_updated(uri); +} + +Task Server::Impl::notify_resource_updated(const std::string& uri) { + auto session = session_snapshot(); + if (!session || session->stopping.load(std::memory_order_acquire)) { co_return; } - auto it = impl_->session->subscriptions.find(uri); - if (it == impl_->session->subscriptions.end() || !it->second) { + + auto owned_uri = std::make_shared(uri); + auto strand = *session->strand; + co_await boost::asio::co_spawn( + strand, notify_resource_updated_on_strand(std::move(session), std::move(owned_uri)), + boost::asio::use_awaitable); +} + +Task Server::Impl::notify_resource_updated_on_strand(std::shared_ptr session, + std::shared_ptr uri) { + if (session->stopping.load(std::memory_order_acquire)) { + co_return; + } + + auto it = session->subscriptions.find(*uri); + if (it == session->subscriptions.end() || !it->second) { co_return; } ResourceUpdatedNotificationParams params; - params.uri = uri; + params.uri = *uri; JSONRPCNotification notification; notification.method = "notifications/resources/updated"; @@ -728,12 +1446,12 @@ Task Server::notify_resource_updated(const std::string& uri) { notification.params = std::move(p); nlohmann::json json_msg = std::move(notification); - co_await impl_->session->transport->write_message(json_msg.dump()); + co_await session->writer->write_message(json_msg.dump()); } -TypeErasedHandler Server::build_middleware_chain(TypeErasedHandler final_handler) { +TypeErasedHandler Server::Impl::build_middleware_chain(TypeErasedHandler final_handler) { auto handler = std::move(final_handler); - for (const auto& mw : std::ranges::reverse_view(impl_->middlewares)) { + for (const auto& mw : std::ranges::reverse_view(middlewares)) { handler = [mw, next = std::move(handler)]( Context& ctx, const nlohmann::json& params) -> Task { co_return co_await mw(ctx, params, next); @@ -742,18 +1460,22 @@ TypeErasedHandler Server::build_middleware_chain(TypeErasedHandler final_handler return handler; } -std::optional Server::paginate(std::size_t total, - const nlohmann::json& json_msg) { - if (impl_->page_size == 0 || total == 0) { +std::optional Server::Impl::paginate(std::size_t total, + const nlohmann::json& json_msg) { + if (page_size == 0 || total == 0) { return PaginationSlice{0, total, std::nullopt}; } std::size_t offset = 0; - if (json_msg.contains("params") && json_msg.at("params").contains("cursor")) { - auto cursor_str = json_msg.at("params").at("cursor").get(); + if (json_msg.contains("params") && detail::has_json_value(json_msg.at("params"), "cursor")) { try { + auto cursor_str = json_msg.at("params").at("cursor").get(); offset = std::stoull(cursor_str); - } catch (...) { + } catch (const nlohmann::json::exception&) { + return std::nullopt; + } catch (const std::invalid_argument&) { + return std::nullopt; + } catch (const std::out_of_range&) { return std::nullopt; } } @@ -762,7 +1484,7 @@ std::optional Server::paginate(std::size_t total, offset = 0; } - auto end = std::min(offset + impl_->page_size, total); + auto end = std::min(offset + page_size, total); std::optional next_cursor; if (end < total) { next_cursor = std::to_string(end); @@ -771,8 +1493,8 @@ std::optional Server::paginate(std::size_t total, return PaginationSlice{offset, end, std::move(next_cursor)}; } -bool Server::has_tool_output_schema(const std::string& name) const { - for (const auto& tool : impl_->tools) { +bool Server::Impl::has_tool_output_schema(const std::string& name) const { + for (const auto& tool : tools) { if (tool.name == name) { return tool.outputSchema.has_value(); } @@ -780,15 +1502,69 @@ bool Server::has_tool_output_schema(const std::string& name) const { return false; } +std::shared_ptr Server::Impl::session_snapshot() const { + std::lock_guard lock(session_mutex); + return session; +} + void Server::reset_session() { - if (impl_ && impl_->session) { - if (impl_->session->transport) { - impl_->session->transport->close(); - } - impl_->session->pending_requests.clear(); - impl_->session->in_flight.clear(); - impl_->session->subscriptions.clear(); - impl_->session.reset(); + if (impl_) { + impl_->reset_session(impl_->session_snapshot()); + } +} + +// Fails every outstanding request on a session and wakes whoever is waiting on it. Must run +// on the session strand: the request maps are plain std::maps that only the strand mutates, +// and the correlation timers belong to that strand. +void Server::Impl::abandon_session_work(const std::shared_ptr& session) { + for (auto& [id, cancelled] : session->in_flight) { + static_cast(id); + cancelled->store(true, std::memory_order_relaxed); + } + for (auto& [id, pending] : session->pending_requests) { + static_cast(id); + if (!pending.completed) { + pending.error = Error{g_CONNECTION_CLOSED, "Server session closed", std::nullopt}; + pending.completed = true; + } + pending.timer->cancel(); + } +} + +void Server::Impl::reset_session(const std::shared_ptr& session) { + if (!session) { + return; + } + + // Runs on whatever thread destroys the Server, so it does only what is safe from any thread: + // `stopping` is atomic, the session pointer is mutex-guarded, and ITransport::close() wakes a + // blocked reader. The request maps belong to the session strand, so abandoning their entries is + // posted there and never waited on: waiting would deadlock on the session strand itself, and + // would never return when the io_context is stopped or was never run. + session->stopping.store(true, std::memory_order_release); + if (session->strand) { + // Guarded for the same reason as the close below: queueing the hand-off allocates, and an + // exception escaping ~Server would terminate the process. + try { + boost::asio::post(*session->strand, [session] { Impl::abandon_session_work(session); }); + } catch (const std::exception&) { + } + } + if (session->transport) { + // Unregistering the session below is what frees the Server for reuse, and it must happen + // even for a transport that cannot close cleanly. This also runs from ~Server, where an + // escaping exception would terminate the process. + try { + session->transport->close(); + } catch (const std::exception&) { + } + } + + { + std::lock_guard lock(session_mutex); + if (this->session == session) { + this->session.reset(); + } } } diff --git a/src/server/server_tool.cpp b/src/server/server_tool.cpp index c74cec6..c0d3791 100644 --- a/src/server/server_tool.cpp +++ b/src/server/server_tool.cpp @@ -10,21 +10,24 @@ void Server::add_tool(const std::string& name, const std::string& description, tool.description = description; tool.inputSchema = input_schema; - register_tool(tool, name, - [h = std::move(handler)](Context& /*ctx*/, - const nlohmann::json& params) -> Task { - try { - co_return h(params); - } catch (const std::exception& e) { - CallToolResult err; - TextContent tc; - tc.text = e.what(); - err.content.emplace_back(std::move(tc)); - err.isError = true; - nlohmann::json j = std::move(err); - co_return j; - } - }); + register_tool( + tool, name, + [h = std::move(handler)](Context& /*ctx*/, const nlohmann::json& params) + -> Task { co_return h(params); }, + detail::ToolResultMode::eNormalize); +} + +void Server::add_raw_tool(const Tool& tool, TypeErasedHandler handler) { + register_tool(tool, tool.name, std::move(handler), detail::ToolResultMode::eValidated); +} + +void Server::add_raw_tool(const Tool& tool, + std::function handler) { + // Exceptions are turned into isError results by the dispatcher, which is also where the + // message is made safe to hand to the peer. + add_raw_tool(tool, + [h = std::move(handler)](Context& /*ctx*/, const nlohmann::json& params) + -> Task { co_return h(params); }); } } // namespace mcp diff --git a/src/transport/http_client.cpp b/src/transport/http_client.cpp index f09edc2..c844862 100644 --- a/src/transport/http_client.cpp +++ b/src/transport/http_client.cpp @@ -14,6 +14,8 @@ #include #include #include +#include +#include #include #include #include @@ -35,15 +37,23 @@ struct HttpClientTransport::Impl { struct SharedState { explicit SharedState(net::strand& strand) - : timer(strand), resolver(strand) {} + : timer(strand), operation_timer(strand), resolver(strand) { + operation_timer.expires_at(std::chrono::steady_clock::time_point::max()); + } net::steady_timer timer; + net::steady_timer operation_timer; net::ip::tcp::resolver resolver; std::optional stream; std::queue queue; std::atomic closed{false}; + mutable std::mutex metadata_mutex; std::string session_id; std::string last_event_id; + std::string protocol_version{std::string(g_LATEST_PROTOCOL_VERSION)}; + std::optional initialize_request_key; + bool operation_active{false}; + bool read_active{false}; }; struct ParsedSseEvent { @@ -51,6 +61,13 @@ struct HttpClientTransport::Impl { std::string id; }; + struct CloseTarget { + std::string host; + std::string port; + std::string path; + std::function bearer_token_provider; + }; + Impl(const net::any_io_executor& executor, const std::string& url) : strand(net::make_strand(executor)), state(std::make_shared(strand)) { auto parsed = parse_url(url); @@ -108,6 +125,12 @@ struct HttpClientTransport::Impl { auto resolved_endpoints = co_await state->resolver.async_resolve(host, port, net::use_awaitable); + // A close() that ran while the resolve could no longer be cancelled found no socket to + // cancel either. From here to the socket opening inside async_connect() nothing suspends, + // so a later close() runs on the strand after it and closes that socket. + if (state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("HttpClientTransport is closed"); + } if (!state->stream) { state->stream.emplace(strand); } @@ -180,29 +203,81 @@ struct HttpClientTransport::Impl { } void enqueue_message(std::string message_text) const { - auto shared_state = state; - net::post(shared_state->timer.get_executor(), - [shared_state, queued_message = std::move(message_text)]() mutable { - shared_state->queue.push(std::move(queued_message)); - shared_state->timer.cancel(); - }); + state->queue.push(std::move(message_text)); + state->timer.cancel(); } void capture_session_id(const StringResponse& response) { auto header_iter = response.find("MCP-Session-Id"); if (header_iter != response.end()) { + std::lock_guard lock(state->metadata_mutex); state->session_id = std::string(header_iter->value()); } } + bool capture_initialize_request(std::string_view message) { + try { + const auto request = nlohmann::json::parse(message); + if (request.is_object() && request.value("method", "") == "initialize" && + request.contains("id") && + (request.at("id").is_string() || request.at("id").is_number_integer())) { + state->initialize_request_key = request.at("id").get().correlation_key(); + return true; + } + } catch (const std::exception&) { + // The protocol layer reports malformed JSON; the transport only + // tracks valid initialize envelopes for header negotiation. + } + return false; + } + + void capture_negotiated_protocol_version(std::string_view message) { + if (!state->initialize_request_key) { + return; + } + + try { + const auto response = nlohmann::json::parse(message); + if (!response.is_object() || !response.contains("id") || + (!response.at("id").is_string() && !response.at("id").is_number_integer()) || + response.at("id").get().correlation_key() != + *state->initialize_request_key) { + return; + } + + state->initialize_request_key.reset(); + if (!response.contains("result") || !response.at("result").is_object()) { + return; + } + const auto& result = response.at("result"); + if (!result.contains("protocolVersion") || !result.at("protocolVersion").is_string()) { + return; + } + + const auto protocol_version = result.at("protocolVersion").get(); + if (is_supported_protocol_version(protocol_version)) { + state->protocol_version = protocol_version; + } + } catch (const std::exception&) { + // The client protocol layer owns response validation. + } + } + void process_response(const StringResponse& response) { if (response.result() == http::status::accepted) { return; } if (response.result_int() >= constants::g_http_bad_request) { - throw std::runtime_error("HTTP request failed with status " + - std::to_string(response.result_int())); + std::string challenge; + const auto challenge_it = response.find(http::field::www_authenticate); + if (challenge_it != response.end()) { + challenge = std::string(challenge_it->value()); + } + throw HttpStatusError( + response.result_int(), + "HTTP request failed with status " + std::to_string(response.result_int()), + std::move(challenge)); } if (response.find(http::field::content_type) == response.end()) { @@ -213,6 +288,7 @@ struct HttpClientTransport::Impl { if (starts_with(content_type_header, "application/json")) { if (!response.body().empty()) { + capture_negotiated_protocol_version(response.body()); enqueue_message(response.body()); } return; @@ -223,8 +299,10 @@ struct HttpClientTransport::Impl { while (!parsed_messages.empty()) { auto& event = parsed_messages.front(); if (!event.id.empty()) { + std::lock_guard lock(state->metadata_mutex); state->last_event_id = event.id; } + capture_negotiated_protocol_version(event.data); enqueue_message(std::move(event.data)); parsed_messages.pop(); } @@ -245,15 +323,146 @@ struct HttpClientTransport::Impl { state->stream.reset(); } + static void complete_operation(const std::shared_ptr& shared_state) { + shared_state->operation_active = false; + boost::system::error_code ignored; + shared_state->operation_timer.cancel(ignored); + } + + static Task run_read(std::shared_ptr impl) { + auto& state = *impl->state; + if (state.read_active) { + throw std::logic_error("HttpClientTransport supports one pending read"); + } + state.read_active = true; + try { + for (;;) { + if (state.closed.load(std::memory_order_acquire)) { + throw std::runtime_error("HttpClientTransport is closed"); + } + + if (!state.queue.empty()) { + auto message_text = std::move(state.queue.front()); + state.queue.pop(); + state.read_active = false; + co_return message_text; + } + + state.timer.expires_at(std::chrono::steady_clock::time_point::max()); + try { + co_await state.timer.async_wait(net::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != net::error::operation_aborted) { + throw; + } + } + } + } catch (...) { + state.read_active = false; + throw; + } + } + + /// Read the provider once, whole, under the mutex that guards it. Callers pin it at the call + /// that starts a request -- write_message() and close(), both on the application thread -- so + /// no code path on the strand reads the member itself afterwards. + std::function pin_bearer_token_provider() const { + std::lock_guard lock(provider_mutex); + return bearer_token_provider; + } + + static Task run_write(std::shared_ptr impl, std::string message, + std::function bearer_token_provider) { + auto& state = *impl->state; + if (state.closed.load(std::memory_order_acquire)) { + throw std::runtime_error("HttpClientTransport is closed"); + } + + while (state.operation_active) { + try { + co_await state.operation_timer.async_wait(net::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != net::error::operation_aborted) { + throw; + } + } + if (state.closed.load(std::memory_order_acquire)) { + throw std::runtime_error("HttpClientTransport is closed"); + } + } + + state.operation_timer.expires_at(std::chrono::steady_clock::time_point::max()); + state.operation_active = true; + const bool initialize_operation = impl->capture_initialize_request(message); + + try { + co_await impl->ensure_connected(); + + StringRequest request{http::verb::post, impl->path, constants::g_http_version_11}; + request.set(http::field::host, impl->host); + request.set(http::field::content_type, "application/json"); + request.set(http::field::accept, "application/json, text/event-stream"); + request.set("MCP-Protocol-Version", state.protocol_version); + if (bearer_token_provider) { + const auto token = bearer_token_provider(); + if (!token.empty()) { + request.set(http::field::authorization, "Bearer " + token); + } + } + if (!state.session_id.empty()) { + request.set("MCP-Session-Id", state.session_id); + } + if (!state.last_event_id.empty()) { + request.set("Last-Event-ID", state.last_event_id); + } + request.body() = std::move(message); + request.prepare_payload(); + + if (!state.stream) { + throw std::runtime_error("HttpClientTransport stream is not initialized"); + } + auto& stream = *state.stream; + stream.expires_after(std::chrono::seconds(constants::g_http_timeout_seconds)); + co_await http::async_write(stream, request, net::use_awaitable); + + beast::flat_buffer response_buffer; + StringResponse response; + co_await http::async_read(stream, response_buffer, response, net::use_awaitable); + + impl->capture_session_id(response); + impl->process_response(response); + if (initialize_operation) { + state.initialize_request_key.reset(); + } + complete_operation(impl->state); + } catch (...) { + if (initialize_operation) { + state.initialize_request_key.reset(); + } + impl->reset_connection(); + complete_operation(impl->state); + // close() closes the socket under the write, and what the socket layer then reports + // depends on the platform and on which operation the close landed in. + if (state.closed.load(std::memory_order_acquire)) { + throw boost::system::system_error(net::error::operation_aborted); + } + throw; + } + } + net::strand strand; std::shared_ptr state; std::string host; std::string port; std::string path; + mutable std::mutex provider_mutex; + /// Guarded by provider_mutex: written by set_bearer_token_provider(), read once per request in + /// pin_bearer_token_provider(). + std::function bearer_token_provider; }; HttpClientTransport::HttpClientTransport(const net::any_io_executor& executor, const std::string& url) - : impl_(std::make_unique(executor, url)) {} + : impl_(std::make_shared(executor, url)) {} HttpClientTransport::~HttpClientTransport() { try { @@ -264,80 +473,36 @@ HttpClientTransport::~HttpClientTransport() { } } -const std::string& HttpClientTransport::session_id() const { return impl_->state->session_id; } - -const std::string& HttpClientTransport::last_event_id() const { return impl_->state->last_event_id; } +std::string HttpClientTransport::session_id() const { + std::lock_guard lock(impl_->state->metadata_mutex); + return impl_->state->session_id; +} -Task HttpClientTransport::read_message() { - auto& state = *impl_->state; - for (;;) { - if (!state.queue.empty()) { - auto message_text = std::move(state.queue.front()); - state.queue.pop(); - co_return message_text; - } +std::string HttpClientTransport::last_event_id() const { + std::lock_guard lock(impl_->state->metadata_mutex); + return impl_->state->last_event_id; +} - if (state.closed.load(std::memory_order_acquire)) { - throw std::runtime_error("HttpClientTransport is closed"); - } +/// Safe to call at any time, including mid-flight: a request already started keeps the provider it +/// pinned, and the next one picks up the new value. +void HttpClientTransport::set_bearer_token_provider(std::function provider) { + std::lock_guard lock(impl_->provider_mutex); + impl_->bearer_token_provider = std::move(provider); +} - state.timer.expires_at(std::chrono::steady_clock::time_point::max()); - try { - co_await state.timer.async_wait(net::use_awaitable); - } catch (const boost::system::system_error& error) { - if (error.code() != net::error::operation_aborted) { - throw; - } - } - } +Task HttpClientTransport::read_message() { + auto impl = impl_; + return net::co_spawn(impl->strand, Impl::run_read(impl), net::use_awaitable); } Task HttpClientTransport::write_message(std::string_view message) { - if (impl_->state->closed.load(std::memory_order_acquire)) { - throw std::runtime_error("HttpClientTransport is closed"); - } - - std::string msg(message); - co_await net::post(impl_->strand, net::use_awaitable); - - if (impl_->state->closed.load(std::memory_order_acquire)) { - throw std::runtime_error("HttpClientTransport is closed"); - } - - try { - co_await impl_->ensure_connected(); - - StringRequest request{http::verb::post, impl_->path, constants::g_http_version_11}; - request.set(http::field::host, impl_->host); - request.set(http::field::content_type, "application/json"); - request.set(http::field::accept, "application/json, text/event-stream"); - request.set("MCP-Protocol-Version", std::string(g_LATEST_PROTOCOL_VERSION)); - if (!impl_->state->session_id.empty()) { - request.set("MCP-Session-Id", impl_->state->session_id); - } - if (!impl_->state->last_event_id.empty()) { - request.set("Last-Event-ID", impl_->state->last_event_id); - } - request.body() = std::move(msg); - request.prepare_payload(); - - if (!impl_->state->stream) { - throw std::runtime_error("HttpClientTransport stream is not initialized"); - } - auto& stream = *impl_->state->stream; - stream.expires_after(std::chrono::seconds(constants::g_http_timeout_seconds)); - co_await http::async_write(stream, request, net::use_awaitable); - - beast::flat_buffer response_buffer; - StringResponse response; - co_await http::async_read(stream, response_buffer, response, net::use_awaitable); - - impl_->capture_session_id(response); - impl_->process_response(response); - } catch (...) { - impl_->reset_connection(); - throw; - } + auto impl = impl_; + // Pinned here, on the caller's thread, before the coroutine is created: this call is where the + // request starts, so this is the value it means. Nothing on the strand reads the member later. + auto bearer_token_provider = impl->pin_bearer_token_provider(); + return net::co_spawn(impl->strand, + Impl::run_write(impl, std::string(message), std::move(bearer_token_provider)), + net::use_awaitable); } void HttpClientTransport::close() { @@ -348,16 +513,43 @@ void HttpClientTransport::close() { auto shared_state = impl_->state; net::post(shared_state->timer.get_executor(), [shared_state]() { shared_state->timer.cancel(); }); - auto active_session_id = shared_state->session_id; + auto close_target = std::make_shared( + Impl::CloseTarget{impl_->host, impl_->port, impl_->path, impl_->pin_bearer_token_provider()}); net::co_spawn( impl_->strand, - [shared_state, active_session_id, host = impl_->host, port = impl_->port, - path = impl_->path]() -> Task { - if (!active_session_id.empty()) { + [shared_state, close_target = std::move(close_target)]() -> Task { + if (shared_state->operation_active) { + shared_state->resolver.cancel(); + if (shared_state->stream) { + // Closed, not cancelled: a cancel reaches only an operation that is pending + // right now, and the write may be between two of them or have one already + // completed and waiting to resume. A closed socket fails the next one too. + beast::error_code ignored; + shared_state->stream->socket().close(ignored); + } + while (shared_state->operation_active) { + try { + co_await shared_state->operation_timer.async_wait(net::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != net::error::operation_aborted) { + co_return; + } + } + } + // A write that had already finished keeps its stream, now with a closed socket. + if (shared_state->stream && !shared_state->stream->socket().is_open()) { + shared_state->stream.reset(); + } + } + + auto active_session_id = std::make_shared(shared_state->session_id); + auto active_protocol_version = + std::make_shared(shared_state->protocol_version); + if (!active_session_id->empty()) { try { if (!shared_state->stream.has_value()) { auto resolved_endpoints = co_await shared_state->resolver.async_resolve( - host, port, net::use_awaitable); + close_target->host, close_target->port, net::use_awaitable); shared_state->stream.emplace(shared_state->timer.get_executor()); shared_state->stream->expires_after( @@ -366,11 +558,17 @@ void HttpClientTransport::close() { net::use_awaitable); } - http::request delete_request{http::verb::delete_, path, - constants::g_http_version_11}; - delete_request.set(http::field::host, host); - delete_request.set("MCP-Session-Id", active_session_id); - delete_request.set("MCP-Protocol-Version", std::string(g_LATEST_PROTOCOL_VERSION)); + http::request delete_request{ + http::verb::delete_, close_target->path, constants::g_http_version_11}; + delete_request.set(http::field::host, close_target->host); + delete_request.set("MCP-Session-Id", *active_session_id); + delete_request.set("MCP-Protocol-Version", *active_protocol_version); + if (close_target->bearer_token_provider) { + const auto token = close_target->bearer_token_provider(); + if (!token.empty()) { + delete_request.set(http::field::authorization, "Bearer " + token); + } + } shared_state->stream->expires_after( std::chrono::seconds(constants::g_http_timeout_seconds)); @@ -394,7 +592,10 @@ void HttpClientTransport::close() { shared_state->stream->socket().close(operation_error); shared_state->stream.reset(); } - shared_state->session_id.clear(); + { + std::lock_guard lock(shared_state->metadata_mutex); + shared_state->session_id.clear(); + } co_return; }, net::detached); diff --git a/src/transport/http_server.cpp b/src/transport/http_server.cpp index c20f914..cbfed2d 100644 --- a/src/transport/http_server.cpp +++ b/src/transport/http_server.cpp @@ -1,3 +1,4 @@ +#include #include #include @@ -12,11 +13,12 @@ #include #include #include +#include #include +#include #include #include #include -#include #include #include #include @@ -31,6 +33,8 @@ namespace beast = boost::beast; namespace http = boost::beast::http; struct HttpServerTransport::Impl { + using Connection = beast::tcp_stream; + struct SharedState { explicit SharedState(boost::asio::strand& execution_strand) : timer(execution_strand) {} @@ -38,6 +42,7 @@ struct HttpServerTransport::Impl { boost::asio::steady_timer timer; std::queue queue; std::atomic closed{false}; + bool read_active{false}; }; struct PendingResponse { @@ -45,6 +50,9 @@ struct HttpServerTransport::Impl { std::optional response_body; std::optional session_header; std::optional event_id; + /// Set only for a request carried under a sentinel id. Holds the id the caller actually + /// sent, so run_write can put it back before the response leaves the transport. + std::optional client_request_id; bool response_ready{false}; }; @@ -111,21 +119,64 @@ struct HttpServerTransport::Impl { request_json.at("method").get() == "initialize"; } + static bool is_discover_request(const nlohmann::json& request_json) { + return request_json.is_object() && request_json.contains("method") && + request_json.at("method").is_string() && + request_json.at("method").get() == "server/discover"; + } + static std::string_view header_value(const StringRequest::const_iterator& header_it) { return {header_it->value().data(), header_it->value().size()}; } - static std::string generate_session_id() { - std::random_device random_device; - std::mt19937 generator(random_device()); - std::uniform_int_distribution distribution(0, constants::g_hex_digits.size() - 1); + static std::string generate_session_id() { return detail::generate_secure_session_id(); } + + void ensure_configurable() const { + if (listening_started || state->closed.load(std::memory_order_acquire)) { + throw std::logic_error("HttpServerTransport configuration must be set before listen()"); + } + } - std::string session_identifier; - session_identifier.reserve(constants::g_session_id_length); - for (std::size_t i = 0; i < constants::g_session_id_length; i++) { - session_identifier.push_back(constants::g_hex_digits[distribution(generator)]); + bool begin_listening() { + std::lock_guard lock(configuration_mutex); + if (state->closed.load(std::memory_order_acquire)) { + return false; } - return session_identifier; + if (listening_started) { + throw std::logic_error("HttpServerTransport::listen() may only be called once"); + } + listening_started = true; + return true; + } + + static void close_connection(const std::shared_ptr& connection) { + boost::system::error_code ignored; + (void)connection->socket().cancel(ignored); + (void)connection->socket().shutdown(boost::asio::ip::tcp::socket::shutdown_both, ignored); + (void)connection->socket().close(ignored); + } + + void close_active_connections() { + for (const auto& connection : active_connections) { + close_connection(connection); + } + } + + /// @brief Render the challenge a 401 will carry, borrowing the metadata URL when unset. + std::string render_challenge(const BearerChallengeConfig& challenge) const { + auto effective = challenge; + if (effective.resource_metadata.empty() && protected_resource_metadata.has_value()) { + effective.resource_metadata = protected_resource_metadata_url(*protected_resource_metadata); + } + return format_www_authenticate(effective); + } + + bool is_unauthenticated_path(const StringRequest& request) const { + if (unauthenticated_paths.empty()) { + return false; + } + const auto path = http_request_path(std::string_view(request.target())); + return unauthenticated_paths.contains(std::string(path)); } bool is_origin_allowed(std::string_view origin_value) const { @@ -136,12 +187,8 @@ struct HttpServerTransport::Impl { } void enqueue_incoming_message(std::string message_payload) const { - auto shared_state = state; - boost::asio::post(shared_state->timer.get_executor(), - [shared_state, payload = std::move(message_payload)]() mutable { - shared_state->queue.push(std::move(payload)); - shared_state->timer.cancel(); - }); + state->queue.push(std::move(message_payload)); + state->timer.cancel(); } static void set_common_headers(StringResponse& response, bool keep_alive) { @@ -204,8 +251,6 @@ struct HttpServerTransport::Impl { } Task validate_post_session(const StringRequest& request) { - co_await boost::asio::post(strand, boost::asio::use_awaitable); - const auto session_header_it = request.find("MCP-Session-Id"); const bool session_header_present = session_header_it != request.end(); @@ -231,8 +276,6 @@ struct HttpServerTransport::Impl { } Task validate_delete_session(const StringRequest& request) { - co_await boost::asio::post(strand, boost::asio::use_awaitable); - if (!session_id.has_value()) { co_return SessionCheckResult{true, {}}; } @@ -250,9 +293,8 @@ struct HttpServerTransport::Impl { } Task>> register_pending_request( - const std::string& request_id_key) { - co_await boost::asio::post(strand, boost::asio::use_awaitable); - + const std::string& request_id_key, + std::optional client_request_id = std::nullopt) { if (pending_responses.contains(request_id_key)) { co_return std::nullopt; } @@ -260,14 +302,24 @@ struct HttpServerTransport::Impl { auto timer_signal = std::make_shared(strand); timer_signal->expires_at(std::chrono::steady_clock::time_point::max()); - pending_responses.emplace(request_id_key, PendingResponse{timer_signal, std::nullopt, - std::nullopt, std::nullopt, false}); + pending_responses.emplace( + request_id_key, PendingResponse{timer_signal, std::nullopt, std::nullopt, std::nullopt, + std::move(client_request_id), false}); co_return timer_signal; } - Task is_response_ready(const std::string& request_id_key) { - co_await boost::asio::post(strand, boost::asio::use_awaitable); + /// A transport-private id for a request that arrives with no session credential. The random + /// prefix is drawn once per transport from the same secure source as the session id, so a + /// peer cannot construct an id that collides with a live sentinel, and the counter keeps + /// concurrent sentinels distinct from each other. + std::string next_sentinel_request_id() { + if (sentinel_id_prefix.empty()) { + sentinel_id_prefix = "mcp-pregate-" + generate_session_id() + "-"; + } + return sentinel_id_prefix + std::to_string(++sentinel_id_counter); + } + Task is_response_ready(const std::string& request_id_key) { const auto pending_it = pending_responses.find(request_id_key); if (pending_it == pending_responses.end()) { co_return true; @@ -277,8 +329,6 @@ struct HttpServerTransport::Impl { } Task consume_pending_response(const std::string& request_id_key) { - co_await boost::asio::post(strand, boost::asio::use_awaitable); - const auto pending_it = pending_responses.find(request_id_key); if (pending_it == pending_responses.end()) { co_return PendingResult{}; @@ -288,11 +338,11 @@ struct HttpServerTransport::Impl { std::move(pending_it->second.session_header), std::move(pending_it->second.event_id)}; pending_responses.erase(pending_it); + sessionless_request_ids.erase(request_id_key); co_return pending_result; } Task terminate_session() { - co_await boost::asio::post(strand, boost::asio::use_awaitable); session_id.reset(); session_active = false; co_return; @@ -313,6 +363,14 @@ struct HttpServerTransport::Impl { "Invalid MCP-Protocol-Version header"); } + if (is_discover_request(request_json)) { + // server/discover advertises protocol versions (see g_DISCOVERABLE_PROTOCOL_VERSIONS) + // outside g_SUPPORTED_PROTOCOL_VERSIONS, and is reachable with zero prior session + // state, so it is exempted from header validation here exactly like initialize is + // exempted from the negotiated-version check below. + return std::nullopt; + } + if (protocol_header_it == request.end() || header_value(protocol_header_it) == negotiated_protocol_version) { return std::nullopt; @@ -333,6 +391,59 @@ struct HttpServerTransport::Impl { return std::nullopt; } + /// @brief Answer a GET for the configured RFC 9728 document, which needs no bearer token. + std::optional serve_protected_resource_metadata( + const StringRequest& request) const { + if (!protected_resource_metadata.has_value() || request.method() != http::verb::get) { + return std::nullopt; + } + if (http_request_path(std::string_view(request.target())) != protected_resource_metadata_path) { + return std::nullopt; + } + if (auto error = check_origin(request)) { + return error; + } + return make_json_response(request, http::status::ok, protected_resource_metadata_body); + } + + StringResponse make_unauthorized_response(const StringRequest& request) const { + auto response = + make_error_response(request, http::status::unauthorized, "Invalid bearer token"); + response.set(http::field::www_authenticate, www_authenticate_value); + return response; + } + + static std::string_view request_bearer_token(const StringRequest& request) { + const auto authorization_it = request.find(http::field::authorization); + return authorization_it == request.end() ? std::string_view{} + : http_bearer_token(header_value(authorization_it)); + } + + /// @brief Run the bearer check when it can be decided without suspending. + /// @details Returns nothing when an async validator is installed, because that decision + /// belongs to check_authorization_async; keeping the two apart leaves the far more + /// common synchronous path free of a coroutine frame and of copying the token. + std::optional check_authorization(const StringRequest& request) const { + if (async_bearer_token_validator || !bearer_token_validator) { + return std::nullopt; + } + + const auto token = request_bearer_token(request); + if (token.empty() || !bearer_token_validator(token)) { + return make_unauthorized_response(request); + } + return std::nullopt; + } + + /// @brief Run the bearer check against an async validator. Only entered when one is installed. + Task> check_authorization_async(const StringRequest& request) const { + const auto token = request_bearer_token(request); + if (token.empty() || !co_await async_bearer_token_validator(std::string(token))) { + co_return make_unauthorized_response(request); + } + co_return std::nullopt; + } + Task handle_post(const StringRequest& request) { const auto request_json = nlohmann::json::parse(request.body(), nullptr, false); if (request_json.is_discarded() || !request_json.is_object()) { @@ -346,11 +457,31 @@ struct HttpServerTransport::Impl { if (auto error = check_origin(request)) { co_return std::move(*error); } + if (auto error = check_authorization(request)) { + co_return std::move(*error); + } + if (async_bearer_token_validator) { + if (auto error = co_await check_authorization_async(request)) { + co_return std::move(*error); + } + } - const auto session_check = co_await validate_post_session(request); - if (!session_check.ok) { - co_return make_error_response(request, http::status::bad_request, - session_check.error_message); + // server/discover is a pre-gate method: it MUST stay reachable with zero prior session state, + // so a discover request with no MCP-Session-Id header skips the POST session gate. Any other + // method, and discover with a session header, still goes through validate_post_session, and + // this path neither creates nor mutates session state. + // + // The "id" requirement keeps a sessionless *notification* named server/discover from skipping + // the gate and pushing its body onto the unbounded incoming queue below. + const bool is_sessionless_discover = is_discover_request(request_json) && + request_json.contains("id") && + request.find("MCP-Session-Id") == request.end(); + if (!is_sessionless_discover) { + const auto session_check = co_await validate_post_session(request); + if (!session_check.ok) { + co_return make_error_response(request, http::status::bad_request, + session_check.error_message); + } } const bool has_request_id = request_json.contains("id"); @@ -360,14 +491,38 @@ struct HttpServerTransport::Impl { co_return make_empty_json_response(request, http::status::accepted); } - const auto request_id_key = request_json.at("id").dump(); - const auto timer_signal = co_await register_pending_request(request_id_key); + // A sessionless discover carries an id chosen by an unauthenticated party, so it is + // registered under a transport-private sentinel id and the caller's own id is restored in + // run_write. The raw id would share a key space with the established session: a prober could + // claim an id the session then needs ("Request id already pending"), and could steer + // replay-store exclusion through sessionless_request_ids. Session-gated requests keep the + // exact bytes the peer sent. + std::string request_id_key; + std::optional client_request_id; + std::string sentinel_body; + if (is_sessionless_discover) { + client_request_id = request_json.at("id"); + auto sentinel_json = request_json; + sentinel_json["id"] = next_sentinel_request_id(); + request_id_key = sentinel_json.at("id").dump(); + sentinel_body = sentinel_json.dump(); + } else { + request_id_key = request_json.at("id").dump(); + } + + const auto timer_signal = co_await register_pending_request(request_id_key, client_request_id); if (!timer_signal.has_value()) { co_return make_error_response(request, http::status::bad_request, "Request id already pending"); } - - enqueue_incoming_message(request.body()); + if (is_sessionless_discover) { + // Tracked so run_write can keep this response out of the shared replay store. The + // entry is dropped again in consume_pending_response and in close(). + sessionless_request_ids.insert(request_id_key); + enqueue_incoming_message(std::move(sentinel_body)); + } else { + enqueue_incoming_message(request.body()); + } for (;;) { if (state->closed.load(std::memory_order_acquire)) { @@ -400,7 +555,8 @@ struct HttpServerTransport::Impl { accept_it != request.end() && std::string_view(accept_it->value()).find("text/event-stream") != std::string_view::npos; - if (!json_only_ && client_accepts_sse && pending_result.event_id.has_value()) { + if (!json_only_.load(std::memory_order_acquire) && client_accepts_sse && + pending_result.event_id.has_value()) { auto response = make_sse_response(request, *pending_result.event_id, *pending_result.response_body); if (pending_result.session_header.has_value()) { @@ -424,6 +580,14 @@ struct HttpServerTransport::Impl { if (auto error = check_origin(request)) { co_return std::move(*error); } + if (auto error = check_authorization(request)) { + co_return std::move(*error); + } + if (async_bearer_token_validator) { + if (auto error = co_await check_authorization_async(request)) { + co_return std::move(*error); + } + } const auto session_check = co_await validate_delete_session(request); if (!session_check.ok) { @@ -440,6 +604,17 @@ struct HttpServerTransport::Impl { if (auto error = check_protocol_version(request, request_json)) { co_return std::move(*error); } + if (auto error = check_origin(request)) { + co_return std::move(*error); + } + if (auto error = check_authorization(request)) { + co_return std::move(*error); + } + if (async_bearer_token_validator) { + if (auto error = co_await check_authorization_async(request)) { + co_return std::move(*error); + } + } if (!session_id.has_value()) { co_return make_error_response(request, http::status::bad_request, "No active session"); @@ -468,6 +643,19 @@ struct HttpServerTransport::Impl { } Task handle_request(const StringRequest& request) { + if (state->closed.load(std::memory_order_acquire)) { + co_return make_error_response(request, http::status::service_unavailable, + "Transport closed"); + } + if (auto document = serve_protected_resource_metadata(request)) { + co_return std::move(*document); + } + if (is_unauthenticated_path(request)) { + // An exempt path is excluded from MCP dispatch, not merely excused from the bearer + // check. Serving MCP here would answer it with no authentication at all, so a path the + // metadata route did not claim has nothing left to answer it. + co_return make_error_response(request, http::status::not_found, "Not found"); + } if (request.method() == http::verb::post) { co_return co_await handle_post(request); } @@ -484,24 +672,59 @@ struct HttpServerTransport::Impl { co_return response; } - Task handle_connection(boost::asio::ip::tcp::socket socket) { - beast::tcp_stream stream(std::move(socket)); + /// @brief Report a body that exceeded max_request_body_bytes, ignoring a dead peer. + static Task write_payload_too_large(Connection& stream) { + StringResponse response{http::status::payload_too_large, 11}; + set_common_headers(response, false); + response.set(http::field::content_type, "application/json"); + response.body() = nlohmann::json{{"error", "Request body too large"}}.dump(); + response.prepare_payload(); + try { + co_await http::async_write(stream, response, boost::asio::use_awaitable); + } catch (const boost::system::system_error&) { + // The peer that overran the limit may already be gone; nothing more to report to it. + (void)0; + } + } + + Task handle_connection(const std::shared_ptr& connection) { + auto& stream = *connection; beast::flat_buffer request_buffer; for (;;) { - StringRequest request; + http::request_parser parser; + parser.body_limit(max_request_body_bytes); + bool body_too_large = false; try { - co_await http::async_read(stream, request_buffer, request, boost::asio::use_awaitable); + co_await http::async_read(stream, request_buffer, parser, boost::asio::use_awaitable); } catch (const boost::system::system_error& err) { if (err.code() == boost::asio::error::eof || err.code() == boost::asio::error::connection_reset || err.code() == boost::asio::error::operation_aborted) { break; } - throw; + if (err.code() != http::error::body_limit) { + throw; + } + body_too_large = true; + } + + if (body_too_large) { + // The parser cannot resynchronize after a truncated body, so the connection ends + // with this answer rather than reading another request from it. + co_await write_payload_too_large(stream); + break; + } + StringRequest request = parser.release(); + + if (state->closed.load(std::memory_order_acquire)) { + break; } auto response = co_await handle_request(request); + if (state->closed.load(std::memory_order_acquire)) { + break; + } const bool keep_connection_alive = response.keep_alive(); co_await http::async_write(stream, response, boost::asio::use_awaitable); @@ -514,27 +737,146 @@ struct HttpServerTransport::Impl { (void)stream.socket().shutdown(boost::asio::ip::tcp::socket::shutdown_send, shutdown_error); } + static Task run_read(std::shared_ptr impl) { + auto& state = *impl->state; + if (state.read_active) { + throw std::logic_error("HttpServerTransport supports one pending read"); + } + state.read_active = true; + + try { + for (;;) { + if (state.closed.load(std::memory_order_acquire)) { + throw std::runtime_error("HttpServerTransport is closed"); + } + + if (!state.queue.empty()) { + auto message_payload = std::move(state.queue.front()); + state.queue.pop(); + state.read_active = false; + co_return message_payload; + } + + state.timer.expires_at(std::chrono::steady_clock::time_point::max()); + try { + co_await state.timer.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + } + } catch (...) { + state.read_active = false; + throw; + } + } + + static Task run_write(std::shared_ptr impl, std::string message) { + if (impl->state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("HttpServerTransport is closed"); + } + + const auto response_json = nlohmann::json::parse(message, nullptr, false); + if (response_json.is_discarded() || !response_json.is_object()) { + co_return; + } + + std::optional response_id_key; + if (response_json.contains("id")) { + response_id_key = response_json.at("id").dump(); + } + + // Responses to sessionless pre-gate requests must never reach the replay store. That + // store is a bounded ring shared with the established session, so appending them would + // evict the session's own replay history and turn its next Last-Event-ID resume into a + // 410 Gone. Sessionless requests are unauthenticated by construction, so their traffic + // must not be able to consume replay capacity that belongs to a real session. + const bool skip_event_store = + response_id_key.has_value() && impl->sessionless_request_ids.contains(*response_id_key); + + std::optional event_id; + if (!impl->json_only_.load(std::memory_order_acquire) && !skip_event_store) { + event_id = impl->event_store.append(message); + } + + if (!response_id_key.has_value()) { + co_return; + } + + const auto& request_id_key = *response_id_key; + const auto pending_it = impl->pending_responses.find(request_id_key); + if (pending_it == impl->pending_responses.end()) { + co_return; + } + + // The request was carried under a sentinel id, so the server answered the sentinel. Put + // the caller's own id back before the response leaves the transport; the peer must see + // the id it sent, and must never see the sentinel. + if (pending_it->second.client_request_id.has_value()) { + auto restored_response = response_json; + restored_response["id"] = *pending_it->second.client_request_id; + message = restored_response.dump(); + } + + pending_it->second.response_body = std::move(message); + pending_it->second.event_id = std::move(event_id); + pending_it->second.response_ready = true; + + if (is_initialize_result_response(response_json)) { + const auto& result = response_json.at("result"); + if (!impl->session_id.has_value()) { + impl->session_id = generate_session_id(); + } + if (result.contains("protocolVersion") && result.at("protocolVersion").is_string()) { + impl->negotiated_protocol_version = result.at("protocolVersion").get(); + } + pending_it->second.session_header = impl->session_id; + impl->session_active = true; + } + + pending_it->second.ready_timer->cancel(); + } + std::string host; unsigned short port; boost::asio::strand strand; boost::asio::ip::tcp::acceptor acceptor; std::shared_ptr state; + std::unordered_set> active_connections; + + mutable std::mutex configuration_mutex; + bool listening_started{false}; std::unordered_map pending_responses; + std::unordered_set sessionless_request_ids; + std::string sentinel_id_prefix; + std::uint64_t sentinel_id_counter{0}; std::optional session_id; std::string negotiated_protocol_version{std::string(g_LATEST_PROTOCOL_VERSION)}; bool session_active{false}; - bool allow_all_origins{true}; + bool allow_all_origins{false}; std::unordered_set allowed_origins; + BearerTokenValidator bearer_token_validator; + AsyncBearerTokenValidator async_bearer_token_validator; + std::size_t max_request_body_bytes{constants::g_default_max_request_body_bytes}; + BearerChallengeConfig bearer_challenge; + /// Rendered once when the challenge or the metadata changes, so serving a 401 never formats. + std::string www_authenticate_value{"Bearer"}; + std::optional protected_resource_metadata; + std::string protected_resource_metadata_body; + /// The resolved serving path: `path` when set, else derived from `resource`. + std::string protected_resource_metadata_path; + std::unordered_set unauthenticated_paths; EventStore event_store; - bool json_only_{false}; + std::atomic json_only_{false}; }; HttpServerTransport::HttpServerTransport(const boost::asio::any_io_executor& executor, std::string host, unsigned short port, std::size_t event_store_capacity) - : impl_(std::make_unique(executor, std::move(host), port, event_store_capacity)) {} + : impl_(std::make_shared(executor, std::move(host), port, event_store_capacity)) {} HttpServerTransport::~HttpServerTransport() { try { @@ -556,126 +898,183 @@ unsigned short HttpServerTransport::port() const { return endpoint.port(); } -void HttpServerTransport::set_json_only(bool json_only) { impl_->json_only_ = json_only; } - -Task HttpServerTransport::read_message() { - auto& state = *impl_->state; - for (;;) { - if (!state.queue.empty()) { - auto message_payload = std::move(state.queue.front()); - state.queue.pop(); - co_return message_payload; - } - - if (state.closed.load(std::memory_order_acquire)) { - throw std::runtime_error("HttpServerTransport is closed"); - } +void HttpServerTransport::set_json_only(bool json_only) { + impl_->json_only_.store(json_only, std::memory_order_release); +} - state.timer.expires_at(std::chrono::steady_clock::time_point::max()); - try { - co_await state.timer.async_wait(boost::asio::use_awaitable); - } catch (const boost::system::system_error& err) { - if (err.code() != boost::asio::error::operation_aborted) { - throw; - } - } +void HttpServerTransport::set_allowed_origins(std::vector origins) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + impl_->allowed_origins.clear(); + impl_->allowed_origins.reserve(origins.size()); + for (auto& origin : origins) { + impl_->allowed_origins.insert(std::move(origin)); } + impl_->allow_all_origins = false; } -Task HttpServerTransport::write_message(std::string_view message) { - std::string msg(message); - co_await boost::asio::post(impl_->strand, boost::asio::use_awaitable); +void HttpServerTransport::set_allow_all_origins(bool allow_all) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + impl_->allow_all_origins = allow_all; +} - if (impl_->state->closed.load(std::memory_order_acquire)) { - throw std::runtime_error("HttpServerTransport is closed"); +void HttpServerTransport::set_bearer_token_validator(BearerTokenValidator validator) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + if (validator && impl_->async_bearer_token_validator) { + throw std::logic_error("HttpServerTransport accepts one bearer token validator"); } + impl_->bearer_token_validator = std::move(validator); +} - const auto response_json = nlohmann::json::parse(msg, nullptr, false); - if (response_json.is_discarded() || !response_json.is_object()) { - co_return; +void HttpServerTransport::set_async_bearer_token_validator(AsyncBearerTokenValidator validator) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + if (validator && impl_->bearer_token_validator) { + throw std::logic_error("HttpServerTransport accepts one bearer token validator"); } + impl_->async_bearer_token_validator = std::move(validator); +} - std::optional event_id; - if (!impl_->json_only_) { - event_id = impl_->event_store.append(msg); +void HttpServerTransport::set_max_request_body_bytes(std::size_t max_bytes) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + if (max_bytes == 0) { + throw std::invalid_argument("Maximum request body size must be greater than zero"); } + impl_->max_request_body_bytes = max_bytes; +} - if (!response_json.contains("id")) { - co_return; +void HttpServerTransport::set_bearer_challenge(BearerChallengeConfig challenge) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + auto rendered = impl_->render_challenge(challenge); + impl_->bearer_challenge = std::move(challenge); + impl_->www_authenticate_value = std::move(rendered); +} + +void HttpServerTransport::set_protected_resource_metadata(ProtectedResourceMetadataConfig metadata) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + if (metadata.resource.empty()) { + throw std::invalid_argument("Protected-resource metadata requires a resource URL"); } - const auto request_id_key = response_json.at("id").dump(); - const auto pending_it = impl_->pending_responses.find(request_id_key); - if (pending_it == impl_->pending_responses.end()) { - co_return; + auto metadata_url = protected_resource_metadata_url(metadata); + auto document = format_protected_resource_metadata(metadata); + auto challenge = impl_->bearer_challenge; + if (challenge.resource_metadata.empty()) { + challenge.resource_metadata = std::move(metadata_url); } + auto rendered = format_www_authenticate(challenge); - pending_it->second.response_body = msg; - pending_it->second.event_id = std::move(event_id); - pending_it->second.response_ready = true; + auto document_path = protected_resource_metadata_path(metadata); - if (Impl::is_initialize_result_response(response_json)) { - const auto& result = response_json.at("result"); - if (!impl_->session_id.has_value()) { - impl_->session_id = Impl::generate_session_id(); - } - if (result.contains("protocolVersion") && result.at("protocolVersion").is_string()) { - impl_->negotiated_protocol_version = result.at("protocolVersion").get(); - } - pending_it->second.session_header = impl_->session_id; - impl_->session_active = true; + impl_->protected_resource_metadata = std::move(metadata); + impl_->protected_resource_metadata_path = std::move(document_path); + impl_->protected_resource_metadata_body = std::move(document); + impl_->www_authenticate_value = std::move(rendered); +} + +void HttpServerTransport::set_unauthenticated_paths(std::vector paths) { + std::lock_guard lock(impl_->configuration_mutex); + impl_->ensure_configurable(); + impl_->unauthenticated_paths.clear(); + impl_->unauthenticated_paths.reserve(paths.size()); + for (auto& path : paths) { + impl_->unauthenticated_paths.insert(std::move(path)); } +} - pending_it->second.ready_timer->cancel(); - co_return; +Task HttpServerTransport::read_message() { + auto impl = impl_; + return boost::asio::co_spawn(impl->strand, Impl::run_read(impl), boost::asio::use_awaitable); +} + +Task HttpServerTransport::write_message(std::string_view message) { + auto impl = impl_; + return boost::asio::co_spawn(impl->strand, Impl::run_write(impl, std::string(message)), + boost::asio::use_awaitable); } void HttpServerTransport::close() { - if (impl_->state->closed.exchange(true, std::memory_order_acq_rel)) { - return; + auto impl = impl_; + { + std::lock_guard lock(impl->configuration_mutex); + if (impl->state->closed.exchange(true, std::memory_order_acq_rel)) { + return; + } } - boost::asio::post(impl_->strand, [this]() { + boost::asio::post(impl->strand, [impl]() { boost::system::error_code ec; - (void)impl_->acceptor.cancel(ec); - (void)impl_->acceptor.close(ec); + (void)impl->acceptor.cancel(ec); + (void)impl->acceptor.close(ec); + impl->close_active_connections(); - for (auto& pending_entry : impl_->pending_responses) { + for (auto& pending_entry : impl->pending_responses) { pending_entry.second.ready_timer->cancel(); } - impl_->pending_responses.clear(); + impl->pending_responses.clear(); + impl->sessionless_request_ids.clear(); - impl_->session_id.reset(); - impl_->negotiated_protocol_version = std::string(g_LATEST_PROTOCOL_VERSION); - impl_->session_active = false; + impl->session_id.reset(); + impl->negotiated_protocol_version = std::string(g_LATEST_PROTOCOL_VERSION); + impl->session_active = false; }); - boost::asio::post(impl_->state->timer.get_executor(), - [state = impl_->state]() { state->timer.cancel(); }); + boost::asio::post(impl->state->timer.get_executor(), + [state = impl->state]() { state->timer.cancel(); }); } Task HttpServerTransport::listen() { + auto impl = impl_; + (void)impl->begin_listening(); + return boost::asio::co_spawn(impl->strand, listen_impl(impl), boost::asio::use_awaitable); +} + +Task HttpServerTransport::listen_impl(std::shared_ptr impl) { for (;;) { - if (impl_->state->closed.load(std::memory_order_acquire)) { + if (impl->state->closed.load(std::memory_order_acquire)) { co_return; } - boost::asio::ip::tcp::socket socket(impl_->strand); + boost::asio::ip::tcp::socket socket(impl->strand); try { - socket = co_await impl_->acceptor.async_accept(boost::asio::use_awaitable); + socket = co_await impl->acceptor.async_accept(boost::asio::use_awaitable); } catch (const boost::system::system_error& err) { - if (impl_->state->closed.load(std::memory_order_acquire) || + if (impl->state->closed.load(std::memory_order_acquire) || err.code() == boost::asio::error::operation_aborted) { co_return; } throw; } - boost::asio::co_spawn(impl_->strand, impl_->handle_connection(std::move(socket)), - [](const std::exception_ptr&) { - // Connection errors (EOF, client disconnect) are normal; - // handled per-connection, not propagated to the accept loop. - }); + if (impl->state->closed.load(std::memory_order_acquire)) { + boost::system::error_code ignored; + (void)socket.close(ignored); + co_return; + } + + auto connection = std::make_shared(std::move(socket)); + impl->active_connections.insert(connection); + + boost::asio::co_spawn( + impl->strand, + [impl, connection]() -> Task { + try { + co_await impl->handle_connection(connection); + } catch (...) { + impl->active_connections.erase(connection); + throw; + } + impl->active_connections.erase(connection); + }, + [](const std::exception_ptr&) { + // Connection errors (EOF, client disconnect) are normal; + // handled per-connection, not propagated to the accept loop. + }); } } diff --git a/src/transport/http_session_manager.cpp b/src/transport/http_session_manager.cpp index 54131bd..f11920b 100644 --- a/src/transport/http_session_manager.cpp +++ b/src/transport/http_session_manager.cpp @@ -1,4 +1,5 @@ #include +#include #include #include @@ -17,16 +18,17 @@ #include #include #include +#include #include #include #include #include -#include #include #include #include #include #include +#include #include #include @@ -44,57 +46,62 @@ namespace detail_session_mgr { struct PendingResponse { std::shared_ptr ready_timer; std::optional response_body; - std::optional event_id; + SseEventList stream_events; bool response_ready{false}; }; struct PendingResult { std::optional response_body; - std::optional event_id; + SseEventList stream_events; }; struct SessionRuntime { std::string session_id; - std::string negotiated_protocol_version{std::string(g_LATEST_PROTOCOL_VERSION)}; - std::shared_ptr client_transport; - MemoryTransport* client_transport_ptr{nullptr}; - Server* server{nullptr}; + std::shared_ptr client_transport; - boost::asio::strand session_strand; + mutable std::mutex protocol_mutex; + std::string negotiated_protocol_version{std::string(g_LATEST_PROTOCOL_VERSION)}; + mutable std::mutex response_mutex; EventStore event_store; std::unordered_map pending_responses; + SseEventList unassigned_stream_events; explicit SessionRuntime( - std::string id, boost::asio::strand strand, - std::size_t event_store_capacity = constants::g_event_store_default_capacity) - : session_id(std::move(id)), session_strand(strand), event_store(event_store_capacity) {} + std::string id, std::size_t event_store_capacity = constants::g_event_store_default_capacity) + : session_id(std::move(id)), event_store(event_store_capacity) {} + + [[nodiscard]] std::string protocol_version() const { + std::lock_guard lock(protocol_mutex); + return negotiated_protocol_version; + } + + void set_protocol_version(std::string version) { + std::lock_guard lock(protocol_mutex); + negotiated_protocol_version = std::move(version); + } }; // ============================================================================ // Helpers // ============================================================================ -inline std::string generate_session_id() { - thread_local std::random_device random_device; - thread_local std::mt19937 generator(random_device()); - thread_local std::uniform_int_distribution distribution( - 0, static_cast(constants::g_hex_digits.size() - 1)); - - std::string session_identifier(constants::g_session_id_length, '0'); - for (char& current_char : session_identifier) { - current_char = constants::g_hex_digits[distribution(generator)]; - } - return session_identifier; -} +inline std::string generate_session_id() { return detail::generate_secure_session_id(); } inline bool is_initialize_request(const nlohmann::json& request_json) { return request_json.is_object() && request_json.contains("method") && + request_json.at("method").is_string() && request_json.at("method").get() == "initialize"; } +inline bool is_discover_request(const nlohmann::json& request_json) { + return request_json.is_object() && request_json.contains("method") && + request_json.at("method").is_string() && + request_json.at("method").get() == "server/discover"; +} + inline std::optional initialize_protocol_version(const nlohmann::json& request_json) { if (!request_json.is_object() || !request_json.contains("params") || !request_json.at("params").is_object()) { @@ -146,21 +153,7 @@ inline StringResponse make_empty_json_response(const StringRequest& request, htt return response; } -inline StringResponse make_sse_response(const StringRequest& request, const std::string& event_id, - const std::string& data) { - std::string sse_body = "id: " + event_id + "\ndata: " + data + "\n\n"; - StringResponse response{http::status::ok, request.version()}; - response.set(http::field::server, "mcp-cpp-sdk"); - response.set(http::field::content_type, "text/event-stream"); - response.set(http::field::cache_control, "no-cache"); - response.keep_alive(request.keep_alive()); - response.body() = std::move(sse_body); - response.prepare_payload(); - return response; -} - -inline StringResponse make_sse_replay_response(const StringRequest& request, - const SseEventList& events) { +inline StringResponse make_sse_response(const StringRequest& request, const SseEventList& events) { std::string sse_body; for (const auto& [id, data] : events) { sse_body += "id: "; @@ -186,6 +179,7 @@ inline StringResponse make_sse_replay_response(const StringRequest& request, // ============================================================================ struct StreamableHttpSessionManager::Impl { + using Connection = beast::tcp_stream; using SessionMap = std::unordered_map>; @@ -193,15 +187,33 @@ struct StreamableHttpSessionManager::Impl { unsigned short port; boost::asio::any_io_executor executor; boost::asio::any_io_executor tool_executor_; + boost::asio::strand listener_strand; boost::asio::ip::tcp::acceptor acceptor; StreamableHttpSessionManager::ServerFactory factory; StreamableHttpSessionManager::CustomRequestHandler custom_handler; std::size_t event_store_capacity; - bool json_only_{false}; + std::shared_ptr> json_only_mode_ = std::make_shared>(false); bool stateless_json_mode_{false}; + bool allow_all_origins_{false}; + std::unordered_set allowed_origins_; + BearerTokenValidator bearer_token_validator_; + AsyncBearerTokenValidator async_bearer_token_validator_; + std::size_t max_request_body_bytes_{constants::g_default_max_request_body_bytes}; + BearerChallengeConfig bearer_challenge_; + /// Rendered once when the challenge or the metadata changes, so serving a 401 never formats. + std::string www_authenticate_value_{"Bearer"}; + std::optional protected_resource_metadata_; + std::string protected_resource_metadata_body_; + /// The resolved serving path: `path` when set, else derived from `resource`. + std::string protected_resource_metadata_path_; + std::unordered_set unauthenticated_paths_; std::atomic closed{false}; + mutable std::mutex configuration_mutex_; + bool listening_started_{false}; + mutable std::mutex connections_mutex_; + std::unordered_set> active_connections_; mutable std::shared_mutex sessions_mutex_; SessionMap sessions; @@ -210,7 +222,8 @@ struct StreamableHttpSessionManager::Impl { : host(std::move(host_arg)), port(port_arg), executor(exec), - acceptor(exec), + listener_strand(boost::asio::make_strand(exec)), + acceptor(listener_strand), factory(std::move(factory_arg)), event_store_capacity(capacity) { boost::system::error_code ec; @@ -241,23 +254,94 @@ struct StreamableHttpSessionManager::Impl { } } + void ensure_configurable() const { + if (listening_started_ || closed.load(std::memory_order_acquire)) { + throw std::logic_error( + "StreamableHttpSessionManager configuration must be set before listen()"); + } + } + + bool begin_listening() { + std::lock_guard lock(configuration_mutex_); + if (closed.load(std::memory_order_acquire)) { + return false; + } + if (listening_started_) { + throw std::logic_error("StreamableHttpSessionManager::listen() may only be called once"); + } + listening_started_ = true; + return true; + } + + static void close_connection_now(const std::shared_ptr& connection) { + boost::system::error_code ignored; + (void)connection->socket().cancel(ignored); + (void)connection->socket().shutdown(boost::asio::ip::tcp::socket::shutdown_both, ignored); + (void)connection->socket().close(ignored); + } + + static void request_connection_close(const std::shared_ptr& connection) { + boost::asio::post(connection->get_executor(), + [connection]() { close_connection_now(connection); }); + } + + bool register_connection(const std::shared_ptr& connection) { + std::lock_guard lock(connections_mutex_); + if (closed.load(std::memory_order_acquire)) { + return false; + } + active_connections_.insert(connection); + return true; + } + + void unregister_connection(const std::shared_ptr& connection) { + std::lock_guard lock(connections_mutex_); + active_connections_.erase(connection); + } + + void close_active_connections() { + std::vector> connections; + { + std::lock_guard lock(connections_mutex_); + connections.reserve(active_connections_.size()); + connections.insert(connections.end(), active_connections_.begin(), + active_connections_.end()); + } + for (const auto& connection : connections) { + request_connection_close(connection); + } + } + // Session lifecycle std::shared_ptr create_session( boost::asio::strand conn_strand) { + if (closed.load(std::memory_order_acquire)) { + return nullptr; + } + auto session_id = detail_session_mgr::generate_session_id(); auto client_mem = std::make_shared(conn_strand); auto server_mem = std::make_shared(conn_strand); client_mem->set_peer(server_mem); server_mem->set_peer(client_mem); - - auto* client_mem_ptr = client_mem.get(); auto server = factory(executor); - Server* server_ptr = server.get(); - std::shared_ptr server_transport = server_mem; - // Use tool_executor_ if set, otherwise fall back to executor + auto runtime = + std::make_shared(session_id, event_store_capacity); + runtime->client_transport = client_mem; + + { + std::unique_lock lock(sessions_mutex_); + if (closed.load(std::memory_order_acquire)) { + client_mem->close(); + return nullptr; + } + sessions.emplace(session_id, runtime); + } + + std::shared_ptr server_transport = server_mem; auto server_exec = tool_executor_ ? tool_executor_ : executor; boost::asio::co_spawn( @@ -268,18 +352,8 @@ struct StreamableHttpSessionManager::Impl { }, boost::asio::detached); - auto runtime = std::make_shared(session_id, conn_strand, - event_store_capacity); - runtime->client_transport = client_mem; - runtime->client_transport_ptr = client_mem_ptr; - runtime->server = server_ptr; - - auto* runtime_ptr = runtime.get(); - (void)runtime_ptr; - { - std::unique_lock lock(sessions_mutex_); - sessions.emplace(session_id, runtime); - } + boost::asio::co_spawn(conn_strand, read_outbound_messages(runtime, json_only_mode_), + boost::asio::detached); return runtime; } @@ -293,44 +367,95 @@ struct StreamableHttpSessionManager::Impl { } void destroy_session(const std::string& session_id) { - std::unique_lock lock(sessions_mutex_); - auto it = sessions.find(session_id); - if (it == sessions.end()) { - return; + std::shared_ptr session; + { + std::unique_lock lock(sessions_mutex_); + auto it = sessions.find(session_id); + if (it == sessions.end()) { + return; + } + session = std::move(it->second); + sessions.erase(it); + } + + std::vector> timers; + { + std::lock_guard lock(session->response_mutex); + timers.reserve(session->pending_responses.size()); + for (auto& [key, pending] : session->pending_responses) { + (void)key; + pending.response_ready = true; + timers.push_back(pending.ready_timer); + } } - for (auto& [key, pending] : it->second->pending_responses) { - (void)pending.ready_timer->cancel(); + for (const auto& timer : timers) { + signal_timer(timer); } - if (it->second->client_transport) { - it->second->client_transport->close(); + if (session->client_transport) { + session->client_transport->close(); } - sessions.erase(it); } // HTTP connection / request handling - Task handle_connection(boost::asio::ip::tcp::socket socket, + /// @brief Report a body that exceeded max_request_body_bytes_, ignoring a dead peer. + static Task write_payload_too_large(Connection& stream) { + StringResponse response{http::status::payload_too_large, 11}; + response.set(http::field::server, "mcp-cpp-sdk"); + response.set(http::field::content_type, "application/json"); + response.keep_alive(false); + response.body() = nlohmann::json{{"error", "Request body too large"}}.dump(); + response.prepare_payload(); + try { + co_await http::async_write(stream, response, boost::asio::use_awaitable); + } catch (const boost::system::system_error&) { + // The peer that overran the limit may already be gone; nothing more to report to it. + (void)0; + } + } + + Task handle_connection(const std::shared_ptr& connection, boost::asio::strand conn_strand) { namespace beast = boost::beast; - beast::tcp_stream stream(std::move(socket)); + auto& stream = *connection; beast::flat_buffer request_buffer; std::unique_ptr stateless_server; for (;;) { - StringRequest request; + http::request_parser parser; + parser.body_limit(max_request_body_bytes_); + bool body_too_large = false; try { - co_await http::async_read(stream, request_buffer, request, boost::asio::use_awaitable); + co_await http::async_read(stream, request_buffer, parser, boost::asio::use_awaitable); } catch (const boost::system::system_error& err) { if (err.code() == boost::asio::error::eof || err.code() == boost::asio::error::connection_reset || err.code() == boost::asio::error::operation_aborted) { break; } - throw; + if (err.code() != http::error::body_limit) { + throw; + } + body_too_large = true; + } + + if (body_too_large) { + // The parser cannot resynchronize after a truncated body, so the connection ends + // with this answer rather than reading another request from it. + co_await write_payload_too_large(stream); + break; + } + StringRequest request = parser.release(); + + if (closed.load(std::memory_order_acquire)) { + break; } auto response = co_await handle_request(request, conn_strand, stateless_server); + if (closed.load(std::memory_order_acquire)) { + break; + } const bool keep_connection_alive = response.keep_alive(); co_await http::async_write(stream, response, boost::asio::use_awaitable); @@ -343,9 +468,79 @@ struct StreamableHttpSessionManager::Impl { (void)stream.socket().shutdown(boost::asio::ip::tcp::socket::shutdown_send, shutdown_error); } + /// @brief Render the challenge a 401 will carry, borrowing the metadata URL when unset. + std::string render_challenge(const BearerChallengeConfig& challenge) const { + auto effective = challenge; + if (effective.resource_metadata.empty() && protected_resource_metadata_.has_value()) { + effective.resource_metadata = + protected_resource_metadata_url(*protected_resource_metadata_); + } + return format_www_authenticate(effective); + } + + bool is_unauthenticated_path(const StringRequest& request) const { + if (unauthenticated_paths_.empty()) { + return false; + } + const auto path = http_request_path(std::string_view(request.target())); + return unauthenticated_paths_.contains(std::string(path)); + } + + /// @brief Answer a GET for the configured RFC 9728 document, which needs no bearer token. + std::optional serve_protected_resource_metadata( + const StringRequest& request) const { + if (!protected_resource_metadata_.has_value() || request.method() != http::verb::get) { + return std::nullopt; + } + if (http_request_path(std::string_view(request.target())) != + protected_resource_metadata_path_) { + return std::nullopt; + } + return detail_session_mgr::make_json_response(request, http::status::ok, + protected_resource_metadata_body_); + } + Task handle_request(const StringRequest& request, boost::asio::strand conn_strand, std::unique_ptr& stateless_server) { + if (closed.load(std::memory_order_acquire)) { + co_return detail_session_mgr::make_error_response( + request, http::status::service_unavailable, "Transport closed"); + } + + const auto origin_it = request.find(http::field::origin); + if (origin_it != request.end() && !allow_all_origins_ && + !allowed_origins_.contains(std::string(origin_it->value()))) { + co_return detail_session_mgr::make_error_response(request, http::status::forbidden, + "Origin not allowed"); + } + + if (auto document = serve_protected_resource_metadata(request)) { + co_return std::move(*document); + } + + const bool unauthenticated_path = is_unauthenticated_path(request); + if (!unauthenticated_path && (bearer_token_validator_ || async_bearer_token_validator_)) { + const auto authorization_it = request.find(http::field::authorization); + const auto token = authorization_it == request.end() + ? std::string_view{} + : http_bearer_token(std::string_view(authorization_it->value())); + // The synchronous path neither copies the token nor suspends; only an async validator + // pays for either. + bool accepted = false; + if (async_bearer_token_validator_) { + accepted = !token.empty() && co_await async_bearer_token_validator_(std::string(token)); + } else if (bearer_token_validator_) { + accepted = !token.empty() && bearer_token_validator_(token); + } + if (!accepted) { + auto response = detail_session_mgr::make_error_response( + request, http::status::unauthorized, "Invalid bearer token"); + response.set(http::field::www_authenticate, www_authenticate_value_); + co_return response; + } + } + if (custom_handler) { auto custom_response = custom_handler(request); if (custom_response.has_value()) { @@ -353,6 +548,15 @@ struct StreamableHttpSessionManager::Impl { } } + if (unauthenticated_path) { + // An exempt path is excluded from MCP dispatch, not merely excused from the bearer + // check. Serving MCP here would answer it with no authentication at all, so once the + // metadata route and the custom handler have both declined, nothing is left to answer + // it. + co_return detail_session_mgr::make_error_response(request, http::status::not_found, + "Not found"); + } + if (request.method() == http::verb::post) { co_return co_await handle_post(request, conn_strand, stateless_server); } @@ -385,11 +589,30 @@ struct StreamableHttpSessionManager::Impl { detail_session_mgr::protocol_header_value(protocol_header_it)); } + /// Dispatching on another executor is spelled as a named coroutine rather than a lambda handed + /// to co_spawn. A closure is built in the caller's frame and carries its captures there, and on + /// GCC 11 an object that lives in a coroutine frame across a suspension can be corrupted -- see + /// the SSO note in docs/contributing.rst. Taking the arguments by value moves them into this + /// coroutine's own frame instead, leaving nothing of the request in the caller's. + static Task dispatch_on_executor(Server* server, nlohmann::json request_json) { + co_return co_await server->dispatch_request_direct(std::move(request_json)); + } + Task handle_stateless_post(const StringRequest& request, nlohmann::json request_json, std::unique_ptr& stateless_server) { + if (closed.load(std::memory_order_acquire)) { + co_return detail_session_mgr::make_error_response( + request, http::status::service_unavailable, "Transport closed"); + } + const bool is_initialize = detail_session_mgr::is_initialize_request(request_json); - if (!is_initialize && !has_valid_protocol_header(request)) { + const bool is_discover = detail_session_mgr::is_discover_request(request_json); + // server/discover advertises protocol versions (see g_DISCOVERABLE_PROTOCOL_VERSIONS) + // outside g_SUPPORTED_PROTOCOL_VERSIONS, which alone backs has_valid_protocol_header; + // exempt it from that check exactly like initialize is exempted, rather than widening + // the check itself for every method. + if (!is_initialize && !is_discover && !has_valid_protocol_header(request)) { co_return detail_session_mgr::make_error_response(request, http::status::bad_request, "Invalid MCP-Protocol-Version header"); } @@ -398,12 +621,35 @@ struct StreamableHttpSessionManager::Impl { co_return detail_session_mgr::make_empty_json_response(request, http::status::accepted); } + const auto dispatch_executor = tool_executor_ ? tool_executor_ : executor; if (!stateless_server) { - stateless_server = factory(executor); + stateless_server = factory(dispatch_executor); } - auto response_body = - co_await stateless_server->dispatch_request_direct(std::move(request_json)); + auto response_body = co_await boost::asio::co_spawn( + dispatch_executor, dispatch_on_executor(stateless_server.get(), std::move(request_json)), + boost::asio::use_awaitable); + co_return detail_session_mgr::make_json_response(request, http::status::ok, + std::move(response_body)); + } + + /// server/discover is a pre-gate method: it MUST be reachable with zero prior state, so a + /// sessionless discover request in stateful mode is dispatched directly against a + /// throwaway Server instance instead of going through resolve_session_for_post — no + /// session is created, registered, or otherwise touched, and no Mcp-Session-Id is issued. + Task handle_sessionless_discover(const StringRequest& request, + nlohmann::json request_json) { + if (closed.load(std::memory_order_acquire)) { + co_return detail_session_mgr::make_error_response( + request, http::status::service_unavailable, "Transport closed"); + } + + const auto dispatch_executor = tool_executor_ ? tool_executor_ : executor; + auto discover_server = factory(dispatch_executor); + + auto response_body = co_await boost::asio::co_spawn( + dispatch_executor, dispatch_on_executor(discover_server.get(), std::move(request_json)), + boost::asio::use_awaitable); co_return detail_session_mgr::make_json_response(request, http::status::ok, std::move(response_body)); } @@ -411,6 +657,11 @@ struct StreamableHttpSessionManager::Impl { std::variant, StringResponse> resolve_session_for_post(const StringRequest& request, const nlohmann::json& request_json, boost::asio::strand conn_strand) { + if (closed.load(std::memory_order_acquire)) { + return detail_session_mgr::make_error_response(request, http::status::service_unavailable, + "Transport closed"); + } + const auto session_header_it = request.find("Mcp-Session-Id"); if (session_header_it == request.end()) { if (!detail_session_mgr::is_initialize_request(request_json)) { @@ -442,13 +693,18 @@ struct StreamableHttpSessionManager::Impl { "Transport closed while waiting response"); } - auto pending_it = session->pending_responses.find(request_id_key); - if (pending_it != session->pending_responses.end() && pending_it->second.response_ready) { - break; + bool response_ready = false; + { + std::lock_guard lock(session->response_mutex); + auto pending_it = session->pending_responses.find(request_id_key); + if (pending_it == session->pending_responses.end()) { + co_return detail_session_mgr::make_error_response( + request, http::status::internal_server_error, "Response lost"); + } + response_ready = pending_it->second.response_ready; } - if (pending_it == session->pending_responses.end()) { - co_return detail_session_mgr::make_error_response( - request, http::status::internal_server_error, "Response lost"); + if (response_ready) { + break; } try { @@ -471,9 +727,9 @@ struct StreamableHttpSessionManager::Impl { accept_it != request.end() && std::string_view(accept_it->value()).find("text/event-stream") != std::string_view::npos; - if (!json_only_ && client_accepts_sse && result.event_id.has_value()) { - auto response = - detail_session_mgr::make_sse_response(request, *result.event_id, *result.response_body); + if (!json_only_mode_->load(std::memory_order_acquire) && client_accepts_sse && + !result.stream_events.empty()) { + auto response = detail_session_mgr::make_sse_response(request, result.stream_events); response.set("Mcp-Session-Id", session->session_id); co_return response; } @@ -487,6 +743,11 @@ struct StreamableHttpSessionManager::Impl { Task handle_post(const StringRequest& request, boost::asio::strand conn_strand, std::unique_ptr& stateless_server) { + if (closed.load(std::memory_order_acquire)) { + co_return detail_session_mgr::make_error_response( + request, http::status::service_unavailable, "Transport closed"); + } + auto request_json = nlohmann::json::parse(request.body(), nullptr, false); if (request_json.is_discarded() || !request_json.is_object()) { co_return detail_session_mgr::make_error_response(request, http::status::bad_request, @@ -507,17 +768,25 @@ struct StreamableHttpSessionManager::Impl { stateless_server); } + if (detail_session_mgr::is_discover_request(request_json) && + request.find("Mcp-Session-Id") == request.end()) { + co_return co_await handle_sessionless_discover(request, std::move(request_json)); + } + auto session_var = resolve_session_for_post(request, request_json, conn_strand); if (std::holds_alternative(session_var)) { co_return std::get(session_var); } auto session = std::get>(session_var); - auto* session_raw = session.get(); + if (session == nullptr || session->client_transport == nullptr) { + co_return detail_session_mgr::make_error_response( + request, http::status::service_unavailable, "Transport closed"); + } if (!is_initialize) { if (protocol_header_it != request.end() && detail_session_mgr::protocol_header_value(protocol_header_it) != - session->negotiated_protocol_version) { + session->protocol_version()) { co_return detail_session_mgr::make_error_response( request, http::status::bad_request, "Invalid MCP-Protocol-Version header"); } @@ -525,48 +794,12 @@ struct StreamableHttpSessionManager::Impl { const auto requested_protocol_version = detail_session_mgr::initialize_protocol_version(request_json) .value_or(std::string(g_LATEST_PROTOCOL_VERSION)); - session->negotiated_protocol_version = - std::string(negotiate_protocol_version(requested_protocol_version)); - } - - if (session == nullptr || session->client_transport_ptr == nullptr) { - co_return detail_session_mgr::make_error_response( - request, http::status::internal_server_error, "Session or transport not found"); - } - - // Fast path for json_only mode: bypass MemoryTransport for tools/call - if (json_only_ && request_json.contains("method") && - request_json.at("method").get() == "tools/call") { - if (!request_json.contains("id")) { - co_return detail_session_mgr::make_error_response(request, http::status::bad_request, - "tools/call requires a request id"); - } - - auto params = request_json.at("params").get(); - - try { - auto result = co_await session->server->invoke_tool(params.name, params.arguments); - JSONRPCResultResponse rpc_response; - rpc_response.id = request_json.at("id").get(); - rpc_response.result = std::move(result); - auto response = detail_session_mgr::make_json_response( - request, http::status::ok, nlohmann::json(std::move(rpc_response)).dump()); - response.set("Mcp-Session-Id", session->session_id); - co_return response; - } catch (const std::exception& e) { - JSONRPCErrorResponse rpc_error; - rpc_error.id = request_json.at("id").get(); - rpc_error.error.code = g_INTERNAL_ERROR; - rpc_error.error.message = e.what(); - auto response = detail_session_mgr::make_json_response( - request, http::status::ok, nlohmann::json(std::move(rpc_error)).dump()); - response.set("Mcp-Session-Id", session->session_id); - co_return response; - } + session->set_protocol_version( + std::string(negotiate_protocol_version(requested_protocol_version))); } if (!request_json.contains("id") || !request_json.contains("method")) { - co_await session->client_transport_ptr->write_json(std::move(request_json)); + co_await session->client_transport->write_json(std::move(request_json)); auto response = detail_session_mgr::make_empty_json_response(request, http::status::accepted); response.set("Mcp-Session-Id", session->session_id); @@ -575,46 +808,23 @@ struct StreamableHttpSessionManager::Impl { const auto request_id_key = request_json.at("id").dump(); - if (session->pending_responses.contains(request_id_key)) { - co_return detail_session_mgr::make_error_response(request, http::status::bad_request, - "Request id already pending"); - } - auto timer_signal = std::make_shared(conn_strand); timer_signal->expires_at(std::chrono::steady_clock::time_point::max()); - session->pending_responses.emplace( - request_id_key, - detail_session_mgr::PendingResponse{timer_signal, std::nullopt, std::nullopt, false}); - - co_await session->client_transport_ptr->write_json(std::move(request_json)); + { + std::lock_guard lock(session->response_mutex); + if (session->pending_responses.contains(request_id_key)) { + co_return detail_session_mgr::make_error_response(request, http::status::bad_request, + "Request id already pending"); + } + session->pending_responses.emplace( + request_id_key, + detail_session_mgr::PendingResponse{timer_signal, std::nullopt, {}, false}); + } - boost::asio::co_spawn( - conn_strand, - [this, session_ptr = session_raw, request_id_key]() -> Task { - try { - auto response_str = co_await session_ptr->client_transport_ptr->read_message(); - boost::asio::post( - session_ptr->session_strand, - [session_ptr, response_str = std::move(response_str), request_id_key]() { - auto it = session_ptr->pending_responses.find(request_id_key); - if (it != session_ptr->pending_responses.end()) { - it->second.response_body = std::move(response_str); - it->second.response_ready = true; - (void)it->second.ready_timer->cancel(); - } - }); - } catch (...) { - boost::asio::post(session_ptr->session_strand, [session_ptr]() { - for (auto& [k, p] : session_ptr->pending_responses) { - (void)p.ready_timer->cancel(); - } - }); - } - }, - boost::asio::detached); + co_await session->client_transport->write_json(std::move(request_json)); - co_return co_await wait_for_response(session_raw, request, request_id_key, timer_signal); + co_return co_await wait_for_response(session.get(), request, request_id_key, timer_signal); } Task handle_delete(const StringRequest& request) { @@ -634,7 +844,7 @@ struct StreamableHttpSessionManager::Impl { const auto protocol_header_it = request.find("MCP-Protocol-Version"); if (protocol_header_it != request.end() && detail_session_mgr::protocol_header_value(protocol_header_it) != - session->negotiated_protocol_version) { + session->protocol_version()) { co_return detail_session_mgr::make_error_response(request, http::status::bad_request, "Invalid MCP-Protocol-Version header"); } @@ -660,7 +870,7 @@ struct StreamableHttpSessionManager::Impl { const auto protocol_header_it = request.find("MCP-Protocol-Version"); if (protocol_header_it != request.end() && detail_session_mgr::protocol_header_value(protocol_header_it) != - session->negotiated_protocol_version) { + session->protocol_version()) { co_return detail_session_mgr::make_error_response(request, http::status::bad_request, "Invalid MCP-Protocol-Version header"); } @@ -668,65 +878,118 @@ struct StreamableHttpSessionManager::Impl { const auto last_event_id_it = request.find("Last-Event-ID"); if (last_event_id_it != request.end()) { auto last_id = std::string(last_event_id_it->value()); - auto missed_events = session->event_store.events_after(last_id); + std::optional missed_events; + { + std::lock_guard lock(session->response_mutex); + missed_events = session->event_store.events_after(last_id); + } if (!missed_events.has_value()) { co_return detail_session_mgr::make_error_response( request, http::status::gone, "Event ID has been evicted from store"); } if (!missed_events->empty()) { - co_return detail_session_mgr::make_sse_replay_response(request, *missed_events); + co_return detail_session_mgr::make_sse_response(request, *missed_events); } } co_return detail_session_mgr::make_empty_json_response(request, http::status::ok); } - static Task dispatch_response_to_pending(detail_session_mgr::SessionRuntime* session, - std::string response_str, bool json_only, - const std::string& request_id_key) { - if (!json_only) { - // Non-json_only mode: parse response and validate id matches - const auto response_json = nlohmann::json::parse(response_str, nullptr, false); - if (response_json.is_discarded() || !response_json.is_object()) { - co_return; - } + static void signal_timer(const std::shared_ptr& timer) { + boost::asio::post(timer->get_executor(), [timer]() { (void)timer->cancel(); }); + } + + static void dispatch_response_to_pending(detail_session_mgr::SessionRuntime* session, + std::string message, bool json_only) { + const auto message_json = nlohmann::json::parse(message, nullptr, false); + if (message_json.is_discarded() || !message_json.is_object()) { + return; + } - auto event_id = session->event_store.append(response_str); + const bool is_response = message_json.contains("id") && + (message_json.contains("result") || message_json.contains("error")); + std::shared_ptr timer_to_signal; - if (!response_json.contains("id") || response_json.at("id").dump() != request_id_key) { - co_return; + { + std::lock_guard lock(session->response_mutex); + + std::optional> stream_event; + if (!json_only) { + auto event_id = session->event_store.append(message); + stream_event.emplace(std::move(event_id), message); } - auto pending_it = session->pending_responses.find(request_id_key); - if (pending_it == session->pending_responses.end()) { - co_return; + if (!is_response) { + if (stream_event.has_value()) { + if (session->pending_responses.size() == 1) { + session->pending_responses.begin()->second.stream_events.push_back( + std::move(*stream_event)); + } else if (!session->pending_responses.empty()) { + session->unassigned_stream_events.push_back(std::move(*stream_event)); + } + } + return; } - pending_it->second.response_body = std::move(response_str); - pending_it->second.event_id = std::move(event_id); - pending_it->second.response_ready = true; - (void)pending_it->second.ready_timer->cancel(); - } else { - // json_only mode: skip parsing, use request_id_key directly + const auto request_id_key = message_json.at("id").dump(); auto pending_it = session->pending_responses.find(request_id_key); if (pending_it == session->pending_responses.end()) { - co_return; + session->unassigned_stream_events.clear(); + return; + } + + if (!json_only) { + auto& events = pending_it->second.stream_events; + events.insert(events.end(), + std::make_move_iterator(session->unassigned_stream_events.begin()), + std::make_move_iterator(session->unassigned_stream_events.end())); + session->unassigned_stream_events.clear(); + events.push_back(std::move(*stream_event)); } - pending_it->second.response_body = std::move(response_str); + pending_it->second.response_body = std::move(message); pending_it->second.response_ready = true; - (void)pending_it->second.ready_timer->cancel(); + timer_to_signal = pending_it->second.ready_timer; + } + + signal_timer(timer_to_signal); + } + + static Task read_outbound_messages( + std::shared_ptr session, + std::shared_ptr> json_only_mode) { + try { + for (;;) { + auto message = co_await session->client_transport->read_message(); + dispatch_response_to_pending(session.get(), std::move(message), + json_only_mode->load(std::memory_order_acquire)); + } + } catch (...) { + std::vector> timers; + { + std::lock_guard lock(session->response_mutex); + timers.reserve(session->pending_responses.size()); + for (auto& [key, pending] : session->pending_responses) { + (void)key; + pending.response_ready = true; + timers.push_back(pending.ready_timer); + } + } + for (const auto& timer : timers) { + signal_timer(timer); + } } } static detail_session_mgr::PendingResult consume_pending_response( detail_session_mgr::SessionRuntime* session, const std::string& request_id_key) { + std::lock_guard lock(session->response_mutex); auto pending_it = session->pending_responses.find(request_id_key); if (pending_it == session->pending_responses.end()) { return {}; } detail_session_mgr::PendingResult result{std::move(pending_it->second.response_body), - std::move(pending_it->second.event_id)}; + std::move(pending_it->second.stream_events)}; session->pending_responses.erase(pending_it); return result; } @@ -740,7 +1003,7 @@ StreamableHttpSessionManager::StreamableHttpSessionManager(const boost::asio::an std::string host, unsigned short port, ServerFactory factory, std::size_t event_store_capacity) - : impl_(std::make_unique(executor, std::move(host), port, std::move(factory), + : impl_(std::make_shared(executor, std::move(host), port, std::move(factory), event_store_capacity)) {} StreamableHttpSessionManager::~StreamableHttpSessionManager() { @@ -753,74 +1016,216 @@ StreamableHttpSessionManager::~StreamableHttpSessionManager() { } void StreamableHttpSessionManager::set_custom_request_handler(CustomRequestHandler handler) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); impl_->custom_handler = std::move(handler); } -std::size_t StreamableHttpSessionManager::session_count() const { return impl_->sessions.size(); } +void StreamableHttpSessionManager::set_allowed_origins(std::vector origins) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + impl_->allowed_origins_.clear(); + impl_->allowed_origins_.reserve(origins.size()); + for (auto& origin : origins) { + impl_->allowed_origins_.insert(std::move(origin)); + } + impl_->allow_all_origins_ = false; +} + +void StreamableHttpSessionManager::set_allow_all_origins(bool allow_all) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + impl_->allow_all_origins_ = allow_all; +} + +void StreamableHttpSessionManager::set_bearer_token_validator(BearerTokenValidator validator) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + if (validator && impl_->async_bearer_token_validator_) { + throw std::logic_error("StreamableHttpSessionManager accepts one bearer token validator"); + } + impl_->bearer_token_validator_ = std::move(validator); +} + +void StreamableHttpSessionManager::set_async_bearer_token_validator( + AsyncBearerTokenValidator validator) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + if (validator && impl_->bearer_token_validator_) { + throw std::logic_error("StreamableHttpSessionManager accepts one bearer token validator"); + } + impl_->async_bearer_token_validator_ = std::move(validator); +} + +void StreamableHttpSessionManager::set_max_request_body_bytes(std::size_t max_bytes) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + if (max_bytes == 0) { + throw std::invalid_argument("Maximum request body size must be greater than zero"); + } + impl_->max_request_body_bytes_ = max_bytes; +} + +void StreamableHttpSessionManager::set_bearer_challenge(BearerChallengeConfig challenge) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + auto rendered = impl_->render_challenge(challenge); + impl_->bearer_challenge_ = std::move(challenge); + impl_->www_authenticate_value_ = std::move(rendered); +} + +void StreamableHttpSessionManager::set_protected_resource_metadata( + ProtectedResourceMetadataConfig metadata) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + if (metadata.resource.empty()) { + throw std::invalid_argument("Protected-resource metadata requires a resource URL"); + } + + auto metadata_url = protected_resource_metadata_url(metadata); + auto document = format_protected_resource_metadata(metadata); + auto challenge = impl_->bearer_challenge_; + if (challenge.resource_metadata.empty()) { + challenge.resource_metadata = std::move(metadata_url); + } + auto rendered = format_www_authenticate(challenge); + + auto document_path = protected_resource_metadata_path(metadata); + + impl_->protected_resource_metadata_ = std::move(metadata); + impl_->protected_resource_metadata_path_ = std::move(document_path); + impl_->protected_resource_metadata_body_ = std::move(document); + impl_->www_authenticate_value_ = std::move(rendered); +} + +void StreamableHttpSessionManager::set_unauthenticated_paths(std::vector paths) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); + impl_->unauthenticated_paths_.clear(); + impl_->unauthenticated_paths_.reserve(paths.size()); + for (auto& path : paths) { + impl_->unauthenticated_paths_.insert(std::move(path)); + } +} + +std::size_t StreamableHttpSessionManager::session_count() const { + std::shared_lock lock(impl_->sessions_mutex_); + return impl_->sessions.size(); +} -void StreamableHttpSessionManager::set_json_only(bool json_only) { impl_->json_only_ = json_only; } +void StreamableHttpSessionManager::set_json_only(bool json_only) { + impl_->json_only_mode_->store(json_only, std::memory_order_release); +} void StreamableHttpSessionManager::set_stateless_json_mode(bool enabled) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); impl_->stateless_json_mode_ = enabled; if (enabled) { - impl_->json_only_ = true; + impl_->json_only_mode_->store(true, std::memory_order_release); } } void StreamableHttpSessionManager::set_tool_executor(const boost::asio::any_io_executor& exec) { + std::lock_guard lock(impl_->configuration_mutex_); + impl_->ensure_configurable(); impl_->tool_executor_ = exec; } void StreamableHttpSessionManager::close() { - if (impl_->closed.exchange(true, std::memory_order_acq_rel)) { - return; + auto impl = impl_; + { + std::lock_guard lock(impl->configuration_mutex_); + if (impl->closed.exchange(true, std::memory_order_acq_rel)) { + return; + } } - boost::asio::post(impl_->executor, [this]() { + boost::asio::post(impl->listener_strand, [impl]() { boost::system::error_code ec; - (void)impl_->acceptor.cancel(ec); - (void)impl_->acceptor.close(ec); + (void)impl->acceptor.cancel(ec); + (void)impl->acceptor.close(ec); + }); + impl->close_active_connections(); + + std::vector> active_sessions; + { + std::unique_lock lock(impl->sessions_mutex_); + active_sessions.reserve(impl->sessions.size()); + for (auto& [id, session] : impl->sessions) { + (void)id; + active_sessions.push_back(std::move(session)); + } + impl->sessions.clear(); + } - std::unique_lock lock(impl_->sessions_mutex_); - for (auto& [id, session] : impl_->sessions) { + for (const auto& session : active_sessions) { + std::vector> timers; + { + std::lock_guard lock(session->response_mutex); + timers.reserve(session->pending_responses.size()); for (auto& [key, pending] : session->pending_responses) { - (void)pending.ready_timer->cancel(); - } - if (session->client_transport) { - session->client_transport->close(); + (void)key; + pending.response_ready = true; + timers.push_back(pending.ready_timer); } } - impl_->sessions.clear(); - }); + for (const auto& timer : timers) { + Impl::signal_timer(timer); + } + if (session->client_transport) { + session->client_transport->close(); + } + } } Task StreamableHttpSessionManager::listen() { + auto impl = impl_; + (void)impl->begin_listening(); + return boost::asio::co_spawn(impl->listener_strand, listen_impl(impl), boost::asio::use_awaitable); +} + +Task StreamableHttpSessionManager::listen_impl(std::shared_ptr impl) { for (;;) { - if (impl_->closed.load(std::memory_order_acquire)) { + if (impl->closed.load(std::memory_order_acquire)) { co_return; } - boost::asio::ip::tcp::socket socket(impl_->executor); + boost::asio::ip::tcp::socket socket(impl->listener_strand); try { - socket = co_await impl_->acceptor.async_accept(boost::asio::use_awaitable); + socket = co_await impl->acceptor.async_accept(boost::asio::use_awaitable); } catch (const boost::system::system_error& err) { - if (impl_->closed.load(std::memory_order_acquire) || + if (impl->closed.load(std::memory_order_acquire) || err.code() == boost::asio::error::operation_aborted) { co_return; } throw; } - auto conn_strand = boost::asio::make_strand(impl_->executor); + auto conn_strand = boost::asio::make_strand(impl->executor); auto native_handle = socket.release(); boost::asio::ip::tcp::socket conn_socket(conn_strand, boost::asio::ip::tcp::v4(), native_handle); - boost::asio::co_spawn(conn_strand, - impl_->handle_connection(std::move(conn_socket), conn_strand), - [](const std::exception_ptr&) { - // Connection errors (EOF, client disconnect) are normal; - // handled per-connection, not propagated to the accept loop. - }); + auto connection = std::make_shared(std::move(conn_socket)); + if (!impl->register_connection(connection)) { + Impl::request_connection_close(connection); + co_return; + } + boost::asio::co_spawn( + conn_strand, + [impl, connection, conn_strand]() mutable -> Task { + try { + co_await impl->handle_connection(connection, conn_strand); + } catch (...) { + impl->unregister_connection(connection); + throw; + } + impl->unregister_connection(connection); + }, + [](const std::exception_ptr&) { + // Connection errors (EOF, client disconnect) are normal; + // handled per-connection, not propagated to the accept loop. + }); } } diff --git a/src/transport/http_types.cpp b/src/transport/http_types.cpp new file mode 100644 index 0000000..b8307c4 --- /dev/null +++ b/src/transport/http_types.cpp @@ -0,0 +1,192 @@ +#include + +#include +#include +#include +#include +#include + +namespace mcp { + +namespace { + +constexpr std::string_view g_well_known_prefix = "/.well-known/oauth-protected-resource"; + +/// @brief The origin and path components of an absolute resource URL. +struct SplitResource { + std::string_view origin; ///< Scheme and authority, with no trailing slash. + std::string_view path; ///< Everything after the authority, empty when there is none. +}; + +/// @brief Split an absolute resource URL, rejecting anything without a scheme and a host. +SplitResource split_resource(std::string_view resource) { + constexpr std::string_view scheme_separator = "://"; + const auto scheme_end = resource.find(scheme_separator); + if (scheme_end == std::string_view::npos) { + throw std::invalid_argument("Protected-resource metadata resource must be an absolute URL"); + } + + const auto authority_start = scheme_end + scheme_separator.size(); + // A query or fragment is not part of a resource identifier; drop it before the split so it + // cannot end up spliced into the metadata path. + const auto trimmed = resource.substr(0, resource.find_first_of("?#", authority_start)); + const auto authority_end = trimmed.find('/', authority_start); + if (authority_end == authority_start || trimmed.size() == authority_start) { + throw std::invalid_argument("Protected-resource metadata resource must name a host"); + } + + if (authority_end == std::string_view::npos) { + return SplitResource{trimmed, {}}; + } + return SplitResource{trimmed.substr(0, authority_end), trimmed.substr(authority_end)}; +} + +/// @brief Reject an explicit metadata path that cannot be appended to an origin as it stands. +/// +/// The path is concatenated onto the origin verbatim, so a relative one does not produce a bad +/// path but a different authority: "https://h" + "evil" resolves to the host "hevil". Precondition: +/// `path` is not empty, which the caller treats as "derive the path" instead. +void validate_metadata_path(std::string_view path) { + if (path.front() != '/') { + throw std::invalid_argument("Protected-resource metadata path must begin with '/'"); + } + + // This function does not normalise, and the result is published to clients as the + // authoritative location of the document. A dot segment is never intentional here, so it is + // refused rather than advertised unresolved. Only a whole segment counts: dots inside a + // segment are ordinary characters, and ".well-known" is the prefix this module derives itself. + for (std::size_t start = 0; start < path.size();) { + auto end = path.find('/', start); + if (end == std::string_view::npos) { + end = path.size(); + } + const auto segment = path.substr(start, end - start); + if (segment == "." || segment == "..") { + throw std::invalid_argument( + "Protected-resource metadata path must not contain a '.' or '..' segment"); + } + start = end + 1; + } +} + +/// @brief True when every byte may appear in an RFC 7235 quoted-string, escaped or not. +bool is_quotable(std::string_view value) { + for (const auto character : value) { + const auto byte = static_cast(character); + if (byte != '\t' && (byte < 0x20 || byte > 0x7E)) { + return false; + } + } + return true; +} + +/// @brief Append `key="value"` with `\` and `"` escaped, separating it from any earlier parameter. +void append_challenge_parameter(std::string& header, bool& first, std::string_view key, + std::string_view value) { + header += first ? " " : ", "; + first = false; + header += key; + header += "=\""; + for (const auto character : value) { + if (character == '\\' || character == '"') { + header += '\\'; + } + header += character; + } + header += '"'; +} + +} // namespace + +std::string_view http_bearer_token(std::string_view authorization_header) { + constexpr std::string_view scheme = "Bearer "; + if (authorization_header.size() <= scheme.size()) { + return {}; + } + + for (std::size_t index = 0; index < scheme.size(); ++index) { + auto actual = authorization_header[index]; + auto expected = scheme[index]; + if (actual >= 'A' && actual <= 'Z') { + actual = static_cast(actual - 'A' + 'a'); + } + if (expected >= 'A' && expected <= 'Z') { + expected = static_cast(expected - 'A' + 'a'); + } + if (actual != expected) { + return {}; + } + } + return authorization_header.substr(scheme.size()); +} + +std::string_view http_request_path(std::string_view target) { + return target.substr(0, target.find_first_of("?#")); +} + +std::string format_www_authenticate(const BearerChallengeConfig& config) { + struct Parameter { + std::string_view key; + std::string_view value; + }; + const Parameter parameters[] = { + {"realm", config.realm}, + {"error", config.error}, + {"scope", config.scope}, + {"resource_metadata", config.resource_metadata}, + }; + + for (const auto& [key, value] : parameters) { + if (!is_quotable(value)) { + throw std::invalid_argument("WWW-Authenticate " + std::string(key) + + " contains a character that cannot be quoted"); + } + } + + std::string header = "Bearer"; + bool first = true; + for (const auto& [key, value] : parameters) { + if (!value.empty()) { + append_challenge_parameter(header, first, key, value); + } + } + return header; +} + +std::string format_protected_resource_metadata(const ProtectedResourceMetadataConfig& metadata) { + nlohmann::json document; + document["resource"] = metadata.resource; + if (!metadata.authorization_servers.empty()) { + document["authorization_servers"] = metadata.authorization_servers; + } + if (!metadata.scopes_supported.empty()) { + document["scopes_supported"] = metadata.scopes_supported; + } + return document.dump(); +} + +std::string protected_resource_metadata_path(const ProtectedResourceMetadataConfig& metadata) { + // Split even when the caller supplied a path, so an unusable resource is reported the same + // way either way. + const auto resource = split_resource(metadata.resource); + if (!metadata.path.empty()) { + validate_metadata_path(metadata.path); + return metadata.path; + } + + // RFC 9728 3.1 inserts the well-known segment between the authority and the resource's own + // path, so a resource at https://host/mcp is described at + // https://host/.well-known/oauth-protected-resource/mcp. + auto resource_path = resource.path; + while (!resource_path.empty() && resource_path.back() == '/') { + resource_path.remove_suffix(1); + } + return std::string(g_well_known_prefix) + std::string(resource_path); +} + +std::string protected_resource_metadata_url(const ProtectedResourceMetadataConfig& metadata) { + return std::string(split_resource(metadata.resource).origin) + + protected_resource_metadata_path(metadata); +} + +} // namespace mcp diff --git a/src/transport/memory.cpp b/src/transport/memory.cpp new file mode 100644 index 0000000..ae36639 --- /dev/null +++ b/src/transport/memory.cpp @@ -0,0 +1,233 @@ +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mcp { + +namespace detail { + +struct MemoryTransportState { + using Message = std::variant; + + struct ReadWaiter { + explicit ReadWaiter(const boost::asio::any_io_executor& executor) : signal(executor) { + signal.expires_at(std::chrono::steady_clock::time_point::max()); + } + + void wake() { + boost::system::error_code ignored; + signal.cancel(ignored); + } + + boost::asio::steady_timer signal; + std::optional message; + }; + + explicit MemoryTransportState(const boost::asio::any_io_executor& executor) + : strand(boost::asio::make_strand(executor)) {} + + ~MemoryTransportState() { close_endpoint(lock_peer()); } + + std::shared_ptr lock_peer() const { + std::lock_guard lock(peer_mutex); + return peer.lock(); + } + + void set_peer(const std::shared_ptr& new_peer) { + std::lock_guard lock(peer_mutex); + peer = new_peer; + } + + static void close_endpoint(const std::shared_ptr& state) { + if (!state || state->closed.exchange(true, std::memory_order_acq_rel)) { + return; + } + + boost::asio::post(state->strand, [state]() { + for (auto& weak_waiter : state->read_waiters) { + if (auto waiter = weak_waiter.lock()) { + waiter->wake(); + } + } + state->read_waiters.clear(); + }); + } + + static std::string as_string(Message message) { + if (auto* raw = std::get_if(&message)) { + return std::move(*raw); + } + return std::get(message).dump(); + } + + static nlohmann::json as_json(Message message) { + if (auto* json = std::get_if(&message)) { + return std::move(*json); + } + return nlohmann::json::parse(std::get(message)); + } + + static Task next_message(std::shared_ptr state) { + if (state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("transport closed"); + } + if (!state->incoming.empty()) { + auto message = std::move(state->incoming.front()); + state->incoming.pop(); + co_return message; + } + + auto waiter = std::make_shared(state->strand); + state->read_waiters.emplace_back(waiter); + while (!waiter->message && !state->closed.load(std::memory_order_acquire)) { + boost::system::error_code error; + co_await waiter->signal.async_wait( + boost::asio::redirect_error(boost::asio::use_awaitable, error)); + if (error && !waiter->message && !state->closed.load(std::memory_order_acquire)) { + throw boost::system::system_error(error); + } + } + + if (state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("transport closed"); + } + co_return std::move(*waiter->message); + } + + void deliver_message(Message message) { + while (!read_waiters.empty()) { + auto waiter = read_waiters.front().lock(); + read_waiters.pop_front(); + if (!waiter) { + continue; + } + + waiter->message.emplace(std::move(message)); + waiter->wake(); + return; + } + incoming.emplace(std::move(message)); + } + + template + static Task deliver(std::shared_ptr sender, + std::shared_ptr peer, MessageValue message) { + if (sender->closed.load(std::memory_order_acquire) || + peer->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("transport closed"); + } + + peer->deliver_message(Message(std::move(message))); + co_return; + } + + template + static Task write(std::shared_ptr state, MessageValue message) { + if (state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("transport closed"); + } + + auto peer = state->lock_peer(); + if (!peer) { + throw std::runtime_error("peer transport not set"); + } + + co_await boost::asio::co_spawn(peer->strand, deliver(state, peer, std::move(message)), + boost::asio::use_awaitable); + } + + boost::asio::strand strand; + std::queue incoming; + std::deque> read_waiters; + std::atomic closed{false}; + + private: + mutable std::mutex peer_mutex; + std::weak_ptr peer; +}; + +} // namespace detail + +namespace { + +Task read_message(std::shared_ptr state) { + co_return detail::MemoryTransportState::as_string( + co_await detail::MemoryTransportState::next_message(std::move(state))); +} + +Task read_json(std::shared_ptr state) { + co_return detail::MemoryTransportState::as_json( + co_await detail::MemoryTransportState::next_message(std::move(state))); +} + +} // namespace + +MemoryTransport::MemoryTransport(const boost::asio::any_io_executor& executor) + : state_(std::make_shared(executor)) {} + +MemoryTransport::~MemoryTransport() = default; + +Task MemoryTransport::read_message() { + auto state = state_; + return boost::asio::co_spawn(state->strand, mcp::read_message(state), boost::asio::use_awaitable); +} + +Task MemoryTransport::read_json() { + auto state = state_; + return boost::asio::co_spawn(state->strand, mcp::read_json(state), boost::asio::use_awaitable); +} + +Task MemoryTransport::write_json(nlohmann::json message) { + auto state = state_; + return boost::asio::co_spawn(state->strand, + detail::MemoryTransportState::write(state, std::move(message)), + boost::asio::use_awaitable); +} + +Task MemoryTransport::write_message(std::string_view message) { + auto state = state_; + return boost::asio::co_spawn(state->strand, + detail::MemoryTransportState::write(state, std::string(message)), + boost::asio::use_awaitable); +} + +void MemoryTransport::close() { + auto state = state_; + auto peer = state->lock_peer(); + detail::MemoryTransportState::close_endpoint(state); + detail::MemoryTransportState::close_endpoint(peer); +} + +void MemoryTransport::set_peer(const std::shared_ptr& peer) { + state_->set_peer(peer ? peer->state_ : nullptr); +} + +std::pair, std::shared_ptr> create_memory_transport_pair( + const boost::asio::any_io_executor& executor) { + auto transport_a = std::make_shared(executor); + auto transport_b = std::make_shared(executor); + transport_a->set_peer(transport_b); + transport_b->set_peer(transport_a); + return {transport_a, transport_b}; +} + +} // namespace mcp diff --git a/src/transport/stdio.cpp b/src/transport/stdio.cpp index aec6b60..e311d67 100644 --- a/src/transport/stdio.cpp +++ b/src/transport/stdio.cpp @@ -1,76 +1,344 @@ #include +#if defined(_WIN32) +#include +#include +#else +#include +#include +#endif + #include +#include #include +#include #include #include #include +#include +#include #include +#include +#include #include +#include +#include #include #include +#include #include #include #include namespace mcp { +namespace { + +constexpr int kStdoutDescriptor = 1; +constexpr int kStderrDescriptor = 2; + +#if defined(_WIN32) +int duplicate_descriptor(int descriptor) { return ::_dup(descriptor); } +int replace_descriptor(int source, int target) { return ::_dup2(source, target); } +int release_descriptor(int descriptor) { return ::_close(descriptor); } +int open_null_device() { return ::_open("NUL", _O_WRONLY); } +std::ptrdiff_t write_descriptor(int descriptor, const char* data, std::size_t size) { + return ::_write(descriptor, data, static_cast(size)); +} +#else +int duplicate_descriptor(int descriptor) { return ::dup(descriptor); } +int replace_descriptor(int source, int target) { return ::dup2(source, target); } +int release_descriptor(int descriptor) { return ::close(descriptor); } +int open_null_device() { return ::open("/dev/null", O_WRONLY); } +std::ptrdiff_t write_descriptor(int descriptor, const char* data, std::size_t size) { + return ::write(descriptor, data, size); +} +#endif + +/// Writes straight through to a file descriptor. The transport flushes after +/// every message, so a buffer here would only add a second place for half a +/// message to sit. +class DescriptorStreambuf final : public std::streambuf { + public: + explicit DescriptorStreambuf(int descriptor) : descriptor_(descriptor) {} + + protected: + std::streamsize xsputn(const char_type* data, std::streamsize size) override { + std::streamsize written = 0; + while (written < size) { + const auto result = + write_descriptor(descriptor_, data + written, static_cast(size - written)); + if (result < 0) { + if (errno == EINTR) { + continue; + } + return written; + } + if (result == 0) { + return written; + } + written += static_cast(result); + } + return written; + } + + int_type overflow(int_type value) override { + if (traits_type::eq_int_type(value, traits_type::eof())) { + return traits_type::not_eof(value); + } + const char_type byte = traits_type::to_char_type(value); + return xsputn(&byte, 1) == 1 ? value : traits_type::eof(); + } + + private: + int descriptor_; +}; + +/// Hands the protocol stream to the transport alone. +/// +/// Duplicating standard output is not enough by itself: a duplicate shares the +/// same open file description, so a stray printf still lands in the same byte +/// stream, and its separate buffer can now flush in the middle of a framed +/// message instead of between two of them. The application's standard output +/// has to point somewhere else, and its diagnostics belong on standard error. +class OwnedStdout { + public: + OwnedStdout() : descriptor_(acquire()), buffer_(descriptor_), stream_(&buffer_) {} + + ~OwnedStdout() { + // Whatever the application buffered belongs on the redirected stream, + // not in the protocol once standard output is handed back. + std::cout.flush(); + std::fflush(stdout); + + replace_descriptor(descriptor_, kStdoutDescriptor); + release_descriptor(descriptor_); + owner_active().store(false, std::memory_order_release); + } + + OwnedStdout(const OwnedStdout&) = delete; + OwnedStdout& operator=(const OwnedStdout&) = delete; + OwnedStdout(OwnedStdout&&) = delete; + OwnedStdout& operator=(OwnedStdout&&) = delete; + + std::ostream& stream() noexcept { return stream_; } + + private: + static std::atomic& owner_active() { + static std::atomic active{false}; + return active; + } + + static int acquire() { + bool unowned = false; + if (!owner_active().compare_exchange_strong(unowned, true, std::memory_order_acq_rel)) { + throw std::runtime_error( + "StdioTransport: another transport already owns this process's standard output"); + } + + // Anything already queued for standard output belongs on the real + // standard output, ahead of the first framed message. + std::cout.flush(); + std::fflush(stdout); + + const int descriptor = duplicate_descriptor(kStdoutDescriptor); + if (descriptor < 0) { + owner_active().store(false, std::memory_order_release); + throw std::runtime_error( + "StdioTransport could not duplicate this process's standard output"); + } + + if (replace_descriptor(kStderrDescriptor, kStdoutDescriptor) < 0) { + // A process started with standard error closed still needs its + // standard output kept off the protocol stream. + const int null_device = open_null_device(); + const bool diverted = + null_device >= 0 && replace_descriptor(null_device, kStdoutDescriptor) >= 0; + if (null_device >= 0) { + release_descriptor(null_device); + } + if (!diverted) { + release_descriptor(descriptor); + owner_active().store(false, std::memory_order_release); + throw std::runtime_error( + "StdioTransport could not redirect this process's standard output away from " + "the protocol stream"); + } + } + + return descriptor; + } + + int descriptor_; + DescriptorStreambuf buffer_; + std::ostream stream_; +}; + +} // namespace struct StdioTransport::Impl { struct SharedState { - explicit SharedState(boost::asio::strand& strand) - : timer(strand) {} + SharedState(const boost::asio::any_io_executor& executor, std::ostream& output) + : output(output), strand(boost::asio::make_strand(executor)), read_signal(strand) { + read_signal.expires_at(std::chrono::steady_clock::time_point::max()); + } + + void wake_reader() noexcept { + boost::system::error_code ignored; + read_signal.cancel(ignored); + } - boost::asio::steady_timer timer; + std::ostream& output; + boost::asio::strand strand; + boost::asio::steady_timer read_signal; std::queue queue; + bool read_pending{false}; + bool input_ended{false}; std::atomic closed{false}; + /// A custom stream buffer can call close() from inside write/flush. + /// Recursive locking keeps that re-entrant close from deadlocking. + std::recursive_mutex output_mutex; }; Impl(const boost::asio::any_io_executor& executor, std::istream& input, std::ostream& output) - : input(input), - output(output), - strand(boost::asio::make_strand(executor)), - state(std::make_shared(strand)) { - state->timer.expires_at(std::chrono::steady_clock::time_point::max()); - } + : input(input), state(std::make_shared(executor, output)) {} + + Impl(const boost::asio::any_io_executor& executor, std::istream& input, OwnStdoutTag) + : owned_output(std::make_unique()), + input(input), + state(std::make_shared(executor, owned_output->stream())) {} ~Impl() { - state->closed.store(true, std::memory_order_release); - if (reader_thread.joinable()) { - reader_thread.join(); + try { + close_state(state); + } catch (...) { + state->closed.store(true, std::memory_order_release); } + join_reader(); } - void ensure_reader_started() { - if (reader_started) { - return; - } - reader_started = true; + std::shared_ptr ensure_reader_started() { auto shared_state = state; - reader_thread = std::thread([&in = input, shared_state]() { - std::string line; - while (std::getline(in, line)) { - if (shared_state->closed.load(std::memory_order_acquire)) { - break; + if (shared_state->closed.load(std::memory_order_acquire)) { + return shared_state; + } + + std::lock_guard lock(reader_mutex); + if (!reader_started && !shared_state->closed.load(std::memory_order_acquire)) { + reader_started = true; + reader_thread = std::thread([&in = input, shared_state]() { + try { + std::string line; + while (std::getline(in, line)) { + if (shared_state->closed.load(std::memory_order_acquire)) { + break; + } + + boost::asio::post(shared_state->strand, + [shared_state, message = std::move(line)]() mutable { + shared_state->queue.push(std::move(message)); + shared_state->wake_reader(); + }); + } + } catch (...) { + // Stream failures are reported to readers as end-of-input. } - boost::asio::post(shared_state->timer.get_executor(), - [shared_state, m = std::move(line)]() mutable { - shared_state->queue.push(std::move(m)); - shared_state->timer.cancel(); - }); + boost::asio::post(shared_state->strand, [shared_state]() { + shared_state->input_ended = true; + shared_state->wake_reader(); + }); + }); + } + return shared_state; + } + + static void close_state(const std::shared_ptr& shared_state) { + const bool first_close = !shared_state->closed.exchange(true, std::memory_order_acq_rel); + + // A write that already acquired the output cannot outlive close(). A + // write queued behind this barrier re-checks closed before touching the + // caller-owned stream. + std::unique_lock output_barrier(shared_state->output_mutex); + output_barrier.unlock(); + + if (first_close) { + boost::asio::post(shared_state->strand, [shared_state]() { shared_state->wake_reader(); }); + } + } + + void join_reader() noexcept { + std::thread reader; + { + std::lock_guard lock(reader_mutex); + reader = std::move(reader_thread); + } + if (reader.joinable()) { + reader.join(); + } + } + + static Task read_on_strand(std::shared_ptr shared_state) { + if (shared_state->read_pending) { + throw std::logic_error("StdioTransport supports only one outstanding read"); + } + + struct ReadReservation { + explicit ReadReservation(SharedState& state) : state(state) { state.read_pending = true; } + ~ReadReservation() { state.read_pending = false; } + + SharedState& state; + } reservation(*shared_state); + + for (;;) { + if (shared_state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("StdioTransport is closed"); } - boost::asio::post(shared_state->timer.get_executor(), [shared_state]() { - shared_state->closed.store(true, std::memory_order_release); - shared_state->timer.cancel(); - }); - }); + if (!shared_state->queue.empty()) { + auto message = std::move(shared_state->queue.front()); + shared_state->queue.pop(); + co_return message; + } + + if (shared_state->input_ended) { + throw std::runtime_error("StdioTransport is closed"); + } + + shared_state->read_signal.expires_at(std::chrono::steady_clock::time_point::max()); + boost::system::error_code error; + co_await shared_state->read_signal.async_wait( + boost::asio::redirect_error(boost::asio::use_awaitable, error)); + if (error && error != boost::asio::error::operation_aborted) { + throw boost::system::system_error(error); + } + } + } + + static Task write_on_strand(std::shared_ptr shared_state, std::string message) { + if (shared_state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("StdioTransport is closed"); + } + + std::lock_guard lock(shared_state->output_mutex); + if (shared_state->closed.load(std::memory_order_acquire)) { + throw std::runtime_error("StdioTransport is closed"); + } + + shared_state->output.write(message.data(), static_cast(message.size())); + shared_state->output.put('\n'); + shared_state->output.flush(); + if (!shared_state->output) { + throw std::runtime_error("StdioTransport failed to write to the output stream"); + } + co_return; } + /// Declared first so it outlives `state`, which holds a reference into it, + /// and so standard output is handed back only after the last write. + std::unique_ptr owned_output; std::istream& input; - std::ostream& output; - boost::asio::strand strand; std::shared_ptr state; + std::mutex reader_mutex; bool reader_started{false}; std::thread reader_thread; }; @@ -79,44 +347,33 @@ StdioTransport::StdioTransport(const boost::asio::any_io_executor& executor, std std::ostream& output) : impl_(std::make_unique(executor, input, output)) {} -StdioTransport::~StdioTransport() = default; +StdioTransport::StdioTransport(const boost::asio::any_io_executor& executor, std::istream& input, + OwnStdoutTag tag) + : impl_(std::make_unique(executor, input, tag)) {} -Task StdioTransport::read_message() { - impl_->ensure_reader_started(); - auto& state = *impl_->state; - for (;;) { - if (!state.queue.empty()) { - auto message = std::move(state.queue.front()); - state.queue.pop(); - co_return message; - } +std::unique_ptr StdioTransport::create_owning_stdout( + const boost::asio::any_io_executor& executor, std::istream& input) { + return std::unique_ptr(new StdioTransport(executor, input, OwnStdoutTag{})); +} - if (state.closed.load(std::memory_order_acquire)) { - throw std::runtime_error("StdioTransport is closed"); - } +StdioTransport::~StdioTransport() = default; - state.timer.expires_at(std::chrono::steady_clock::time_point::max()); - try { - co_await state.timer.async_wait(boost::asio::use_awaitable); - } catch (const boost::system::system_error& err) { - if (err.code() != boost::asio::error::operation_aborted) { - throw; - } - } - } +Task StdioTransport::read_message() { + auto state = impl_->ensure_reader_started(); + return boost::asio::co_spawn(state->strand, Impl::read_on_strand(state), + boost::asio::use_awaitable); } Task StdioTransport::write_message(std::string_view message) { - std::string msg(message); - co_await boost::asio::post(impl_->strand, boost::asio::use_awaitable); - impl_->output << msg << '\n'; - impl_->output.flush(); - co_return; + auto state = impl_->state; + auto owned_message = std::string(message); + return boost::asio::co_spawn(state->strand, Impl::write_on_strand(state, std::move(owned_message)), + boost::asio::use_awaitable); } void StdioTransport::close() { - impl_->state->closed.store(true, std::memory_order_release); - impl_->state->timer.cancel(); + auto state = impl_->state; + Impl::close_state(state); } } // namespace mcp diff --git a/src/transport/websocket.cpp b/src/transport/websocket.cpp index e4900c9..3e5eaef 100644 --- a/src/transport/websocket.cpp +++ b/src/transport/websocket.cpp @@ -1,15 +1,21 @@ #include +#include #include #include #include +#include +#include +#include #include #include #include #include #include #include +#include #include +#include #include #include #include @@ -24,30 +30,256 @@ namespace beast = boost::beast; namespace asio = boost::asio; using WsStream = beast::websocket::stream; +namespace { + +constexpr auto kNever = std::chrono::steady_clock::time_point::max(); + +struct OperationWaiter { + explicit OperationWaiter(const asio::any_io_executor& executor) : signal(executor) { + signal.expires_at(kNever); + } + + void wake() noexcept { + boost::system::error_code ignored; + signal.cancel(ignored); + } + + asio::steady_timer signal; + bool granted{false}; +}; + +template +void wake_all(WaiterContainer& waiters) noexcept { + for (const auto& waiter : waiters) { + waiter->wake(); + } + waiters.clear(); +} + +template +void remove_waiter(WaiterContainer& waiters, const std::shared_ptr& waiter) noexcept { + auto position = std::find(waiters.begin(), waiters.end(), waiter); + if (position != waiters.end()) { + waiters.erase(position); + } +} + +void require_open(const std::atomic& closed, std::string_view transport_name) { + if (closed.load(std::memory_order_acquire)) { + throw std::runtime_error(std::string(transport_name) + " is closed"); + } +} + +std::exception_ptr closed_error(std::string_view transport_name) { + return std::make_exception_ptr(std::runtime_error(std::string(transport_name) + " is closed")); +} + +void close_socket(WsStream& ws) noexcept { + beast::error_code ignored; + auto& socket = ws.next_layer().socket(); + socket.cancel(ignored); + socket.shutdown(asio::ip::tcp::socket::shutdown_both, ignored); + socket.close(ignored); +} + +class ReadReservation { + public: + explicit ReadReservation(bool& active) : active_(active) { active_ = true; } + ~ReadReservation() { active_ = false; } + + ReadReservation(const ReadReservation&) = delete; + ReadReservation& operator=(const ReadReservation&) = delete; + + private: + bool& active_; +}; + +class SerializedWriteGate { + public: + SerializedWriteGate(const asio::any_io_executor& executor, const std::atomic& closed, + std::string_view transport_name) + : executor_(executor), closed_(closed), transport_name_(transport_name) {} + + Task acquire() { + require_open(closed_, transport_name_); + if (!active_) { + active_ = true; + co_return; + } + + auto waiter = std::make_shared(executor_); + waiters_.push_back(waiter); + boost::system::error_code error; + co_await waiter->signal.async_wait(asio::redirect_error(asio::use_awaitable, error)); + if (!waiter->granted) { + remove_waiter(waiters_, waiter); + } + if (error && error != asio::error::operation_aborted) { + if (waiter->granted) { + release(); + } + throw boost::system::system_error(error); + } + if (!waiter->granted) { + require_open(closed_, transport_name_); + throw std::runtime_error(transport_name_ + " write was interrupted"); + } + if (closed_.load(std::memory_order_acquire)) { + release(); + require_open(closed_, transport_name_); + } + } + + void release() noexcept { + if (closed_.load(std::memory_order_acquire) || waiters_.empty()) { + active_ = false; + return; + } + + auto waiter = std::move(waiters_.front()); + waiters_.pop_front(); + waiter->granted = true; + waiter->wake(); + } + + void cancel() noexcept { wake_all(waiters_); } + + private: + asio::any_io_executor executor_; + const std::atomic& closed_; + std::string transport_name_; + std::deque> waiters_; + bool active_{false}; +}; + +} // namespace + // ============================================================================ // WebSocketServerTransport::Impl // ============================================================================ struct WebSocketServerTransport::Impl { + enum class HandshakeState { + Pending, + Accepting, + Open, + Failed, + }; + WsStream ws; asio::strand strand; std::atomic closed{false}; - bool handshake_done{false}; + HandshakeState handshake_state{HandshakeState::Pending}; + std::exception_ptr handshake_error; + std::vector> handshake_waiters; + SerializedWriteGate write_gate; + bool read_active{false}; explicit Impl(asio::ip::tcp::socket socket) - : ws(std::move(socket)), strand(asio::make_strand(ws.get_executor())) { + : ws(std::move(socket)), + strand(asio::make_strand(ws.get_executor())), + write_gate(strand, closed, "WebSocketServerTransport") { ws.text(true); } + + static void throw_if_closed(const std::shared_ptr& state) { + require_open(state->closed, "WebSocketServerTransport"); + } + + static void notify_handshake_waiters(const std::shared_ptr& state) { + wake_all(state->handshake_waiters); + } + + static Task wait_for_handshake(std::shared_ptr state) { + auto waiter = std::make_shared(state->strand); + state->handshake_waiters.push_back(waiter); + + boost::system::error_code error; + co_await waiter->signal.async_wait(asio::redirect_error(asio::use_awaitable, error)); + remove_waiter(state->handshake_waiters, waiter); + if (error && error != asio::error::operation_aborted) { + throw boost::system::system_error(error); + } + + throw_if_closed(state); + if (state->handshake_state == HandshakeState::Open) { + co_return; + } + if (state->handshake_error) { + std::rethrow_exception(state->handshake_error); + } + throw std::runtime_error("WebSocket server handshake did not complete"); + } + + static Task ensure_handshake(std::shared_ptr state) { + throw_if_closed(state); + if (state->handshake_state == HandshakeState::Open) { + co_return; + } + if (state->handshake_state == HandshakeState::Accepting) { + co_await wait_for_handshake(std::move(state)); + co_return; + } + if (state->handshake_state == HandshakeState::Failed) { + std::rethrow_exception(state->handshake_error); + } + + state->handshake_state = HandshakeState::Accepting; + try { + co_await state->ws.async_accept(asio::use_awaitable); + throw_if_closed(state); + state->handshake_state = HandshakeState::Open; + } catch (...) { + state->handshake_error = std::current_exception(); + state->handshake_state = HandshakeState::Failed; + notify_handshake_waiters(state); + throw; + } + notify_handshake_waiters(state); + } + + static Task read(std::shared_ptr state) { + throw_if_closed(state); + if (state->read_active) { + throw std::logic_error("WebSocketServerTransport supports only one outstanding read"); + } + + ReadReservation reservation(state->read_active); + + co_await ensure_handshake(state); + beast::flat_buffer buffer; + co_await state->ws.async_read(buffer, asio::use_awaitable); + co_return beast::buffers_to_string(buffer.data()); + } + + static Task write(std::shared_ptr state, std::string message) { + throw_if_closed(state); + co_await ensure_handshake(state); + co_await state->write_gate.acquire(); + try { + co_await state->ws.async_write(asio::buffer(message), asio::use_awaitable); + } catch (...) { + state->write_gate.release(); + throw; + } + state->write_gate.release(); + } + + static void close_on_strand(const std::shared_ptr& state) { + state->handshake_error = closed_error("WebSocketServerTransport"); + state->handshake_state = HandshakeState::Failed; + notify_handshake_waiters(state); + state->write_gate.cancel(); + close_socket(state->ws); + } }; WebSocketServerTransport::WebSocketServerTransport(asio::ip::tcp::socket socket) - : impl_(std::make_unique(std::move(socket))) {} + : impl_(std::make_shared(std::move(socket))) {} WebSocketServerTransport::~WebSocketServerTransport() { try { - if (!impl_->closed.load(std::memory_order_acquire)) { - close(); - } + close(); } catch (...) { // Swallow exceptions in destructor to prevent std::terminate. (void)0; @@ -55,42 +287,23 @@ WebSocketServerTransport::~WebSocketServerTransport() { } Task WebSocketServerTransport::read_message() { - if (impl_->closed.load(std::memory_order_acquire)) { - throw std::runtime_error("WebSocketServerTransport is closed"); - } - - if (!impl_->handshake_done) { - co_await impl_->ws.async_accept(asio::use_awaitable); - impl_->handshake_done = true; - } - - beast::flat_buffer buffer; - co_await impl_->ws.async_read(buffer, asio::use_awaitable); - co_return beast::buffers_to_string(buffer.data()); + auto state = impl_; + return asio::co_spawn(state->strand, Impl::read(state), asio::use_awaitable); } Task WebSocketServerTransport::write_message(std::string_view message) { - if (impl_->closed.load(std::memory_order_acquire)) { - throw std::runtime_error("WebSocketServerTransport is closed"); - } - - if (!impl_->handshake_done) { - co_await impl_->ws.async_accept(asio::use_awaitable); - impl_->handshake_done = true; - } - - co_await impl_->ws.async_write(asio::buffer(message.data(), message.size()), asio::use_awaitable); + auto state = impl_; + auto owned_message = std::string(message); + return asio::co_spawn(state->strand, Impl::write(state, std::move(owned_message)), + asio::use_awaitable); } void WebSocketServerTransport::close() { - if (impl_->closed.exchange(true, std::memory_order_acq_rel)) { + auto state = impl_; + if (state->closed.exchange(true, std::memory_order_acq_rel)) { return; } - - beast::error_code ec; - auto& socket = impl_->ws.next_layer().socket(); - (void)socket.shutdown(asio::ip::tcp::socket::shutdown_both, ec); - (void)socket.close(ec); + asio::post(state->strand, [state]() { Impl::close_on_strand(state); }); } // ============================================================================ @@ -102,8 +315,10 @@ struct WebSocketClientTransport::Impl { Disconnected, Connecting, Connected, + Failed, }; + std::chrono::milliseconds connect_timeout; asio::strand strand; asio::ip::tcp::resolver resolver; WsStream ws; @@ -111,108 +326,141 @@ struct WebSocketClientTransport::Impl { std::string port; std::string path; std::atomic closed{false}; - std::atomic connection_state{ConnectionState::Disconnected}; + ConnectionState connection_state{ConnectionState::Disconnected}; std::exception_ptr connect_error; - std::vector> connection_waiters; + std::vector> connection_waiters; + SerializedWriteGate write_gate; + bool read_active{false}; Impl(const asio::any_io_executor& executor, std::string host_arg, std::string port_arg, - std::string path_arg) - : strand(asio::make_strand(executor)), + std::string path_arg, std::chrono::milliseconds connect_timeout_arg) + : connect_timeout(connect_timeout_arg), + strand(asio::make_strand(executor)), resolver(strand), ws(strand), host(std::move(host_arg)), port(std::move(port_arg)), - path(std::move(path_arg)) { + path(std::move(path_arg)), + write_gate(strand, closed, "WebSocketClientTransport") { ws.text(true); } - void remove_connection_waiter(const std::shared_ptr& waiter) { - for (auto it = connection_waiters.begin(); it != connection_waiters.end(); ++it) { - if (*it == waiter) { - connection_waiters.erase(it); - break; - } - } + static void throw_if_closed(const std::shared_ptr& state) { + require_open(state->closed, "WebSocketClientTransport"); } - void notify_connection_waiters() { - for (const auto& waiter : connection_waiters) { - waiter->cancel(); - } - connection_waiters.clear(); + static void notify_connection_waiters(const std::shared_ptr& state) { + wake_all(state->connection_waiters); } - Task wait_for_connection() { - auto waiter = std::make_shared(strand); - waiter->expires_at(std::chrono::steady_clock::time_point::max()); - connection_waiters.push_back(waiter); - - try { - co_await waiter->async_wait(asio::use_awaitable); - } catch (const boost::system::system_error& err) { - remove_connection_waiter(waiter); - if (err.code() != asio::error::operation_aborted) { - throw; - } + static Task wait_for_connection(std::shared_ptr state) { + auto waiter = std::make_shared(state->strand); + state->connection_waiters.push_back(waiter); + boost::system::error_code error; + co_await waiter->signal.async_wait(asio::redirect_error(asio::use_awaitable, error)); + remove_waiter(state->connection_waiters, waiter); + if (error && error != asio::error::operation_aborted) { + throw boost::system::system_error(error); } - remove_connection_waiter(waiter); - - if (connection_state.load(std::memory_order_acquire) == ConnectionState::Connected) { + throw_if_closed(state); + if (state->connection_state == ConnectionState::Connected) { co_return; } - - if (connect_error) { - std::rethrow_exception(connect_error); + if (state->connect_error) { + std::rethrow_exception(state->connect_error); } - throw std::runtime_error("WebSocket connection did not complete"); } - Task ensure_connected() { - if (connection_state.load(std::memory_order_acquire) == ConnectionState::Connected) { + static Task ensure_connected(std::shared_ptr state) { + throw_if_closed(state); + if (state->connection_state == ConnectionState::Connected) { co_return; } - - if (connection_state.load(std::memory_order_acquire) == ConnectionState::Connecting) { - co_await wait_for_connection(); + if (state->connection_state == ConnectionState::Connecting) { + co_await wait_for_connection(std::move(state)); co_return; } + if (state->connection_state == ConnectionState::Failed) { + std::rethrow_exception(state->connect_error); + } - connection_state.store(ConnectionState::Connecting, std::memory_order_release); - connect_error = nullptr; - + state->connection_state = ConnectionState::Connecting; + state->connect_error = nullptr; try { - auto results = co_await resolver.async_resolve(host, port, asio::use_awaitable); - co_await ws.next_layer().async_connect(*results.begin(), asio::use_awaitable); - co_await ws.async_handshake(host + ":" + port, path, asio::use_awaitable); + auto results = + co_await state->resolver.async_resolve(state->host, state->port, asio::use_awaitable); + // A close() that ran while the resolve could no longer be cancelled closed a socket + // that was not open yet. From here to the socket opening inside async_connect() + // nothing suspends, so a later close() runs on the strand after it and closes it. + throw_if_closed(state); + // One deadline covers the TCP connect and the handshake: a peer that accepts and then + // says nothing would otherwise hold the pending call forever. + if (state->connect_timeout > std::chrono::milliseconds::zero()) { + state->ws.next_layer().expires_after(state->connect_timeout); + } + co_await state->ws.next_layer().async_connect(results, asio::use_awaitable); + co_await state->ws.async_handshake(state->host + ":" + state->port, state->path, + asio::use_awaitable); + state->ws.next_layer().expires_never(); + throw_if_closed(state); + state->connection_state = ConnectionState::Connected; + } catch (...) { + state->connect_error = std::current_exception(); + state->connection_state = ConnectionState::Failed; + notify_connection_waiters(state); + throw; + } + notify_connection_waiters(state); + } + + static Task read(std::shared_ptr state) { + throw_if_closed(state); + if (state->read_active) { + throw std::logic_error("WebSocketClientTransport supports only one outstanding read"); + } + + ReadReservation reservation(state->read_active); + + co_await ensure_connected(state); + beast::flat_buffer buffer; + co_await state->ws.async_read(buffer, asio::use_awaitable); + co_return beast::buffers_to_string(buffer.data()); + } - connection_state.store(ConnectionState::Connected, std::memory_order_release); - notify_connection_waiters(); + static Task write(std::shared_ptr state, std::string message) { + throw_if_closed(state); + co_await ensure_connected(state); + co_await state->write_gate.acquire(); + try { + co_await state->ws.async_write(asio::buffer(message), asio::use_awaitable); } catch (...) { - connect_error = std::current_exception(); - connection_state.store(ConnectionState::Disconnected, std::memory_order_release); - notify_connection_waiters(); - std::rethrow_exception(connect_error); + state->write_gate.release(); + throw; } + state->write_gate.release(); + } + + static void close_on_strand(const std::shared_ptr& state) { + state->connect_error = closed_error("WebSocketClientTransport"); + state->connection_state = ConnectionState::Failed; + notify_connection_waiters(state); + state->write_gate.cancel(); + state->resolver.cancel(); + close_socket(state->ws); } }; WebSocketClientTransport::WebSocketClientTransport(const asio::any_io_executor& executor, - std::string host, std::string port, std::string path) - : impl_(std::make_unique(executor, std::move(host), std::move(port), std::move(path))) {} + std::string host, std::string port, std::string path, + std::chrono::milliseconds connect_timeout) + : impl_(std::make_shared(executor, std::move(host), std::move(port), std::move(path), + connect_timeout)) {} WebSocketClientTransport::~WebSocketClientTransport() { try { - impl_->closed.store(true, std::memory_order_release); - impl_->connection_state.store(Impl::ConnectionState::Disconnected, std::memory_order_release); - impl_->notify_connection_waiters(); - if (impl_->ws.next_layer().socket().is_open()) { - beast::error_code ec; - auto& socket = impl_->ws.next_layer().socket(); - (void)socket.shutdown(asio::ip::tcp::socket::shutdown_both, ec); - (void)socket.close(ec); - } + close(); } catch (...) { // Swallow exceptions in destructor to prevent std::terminate. (void)0; @@ -220,45 +468,23 @@ WebSocketClientTransport::~WebSocketClientTransport() { } Task WebSocketClientTransport::read_message() { - if (impl_->closed.load(std::memory_order_acquire)) { - throw std::runtime_error("WebSocketClientTransport is closed"); - } - - co_await impl_->ensure_connected(); - - beast::flat_buffer buffer; - co_await impl_->ws.async_read(buffer, asio::use_awaitable); - co_return beast::buffers_to_string(buffer.data()); + auto state = impl_; + return asio::co_spawn(state->strand, Impl::read(state), asio::use_awaitable); } Task WebSocketClientTransport::write_message(std::string_view message) { - if (impl_->closed.load(std::memory_order_acquire)) { - throw std::runtime_error("WebSocketClientTransport is closed"); - } - - co_await impl_->ensure_connected(); - - co_await impl_->ws.async_write(asio::buffer(message.data(), message.size()), asio::use_awaitable); + auto state = impl_; + auto owned_message = std::string(message); + return asio::co_spawn(state->strand, Impl::write(state, std::move(owned_message)), + asio::use_awaitable); } void WebSocketClientTransport::close() { - if (impl_->closed.exchange(true, std::memory_order_acq_rel)) { + auto state = impl_; + if (state->closed.exchange(true, std::memory_order_acq_rel)) { return; } - - impl_->connect_error = - std::make_exception_ptr(std::runtime_error("WebSocketClientTransport is closed")); - impl_->connection_state.store(Impl::ConnectionState::Disconnected, std::memory_order_release); - impl_->notify_connection_waiters(); - - if (!impl_->ws.next_layer().socket().is_open()) { - return; - } - - beast::error_code ec; - auto& socket = impl_->ws.next_layer().socket(); - (void)socket.shutdown(asio::ip::tcp::socket::shutdown_both, ec); - (void)socket.close(ec); + asio::post(state->strand, [state]() { Impl::close_on_strand(state); }); } } // namespace mcp diff --git a/test/auth/auth_authorization_test.cpp b/test/auth/auth_authorization_test.cpp new file mode 100644 index 0000000..61ddff6 --- /dev/null +++ b/test/auth/auth_authorization_test.cpp @@ -0,0 +1,4753 @@ +/** + * @file auth_authorization_test.cpp + * @brief Loopback tests for challenge-driven OAuth authorization and the outbound-request controls + * that guard it. + * + * These exercise the full challenge -> discovery -> authorize -> token -> replay contract against + * loopback fixtures. + */ + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +// Reach-in for retained-state assertions; see the header. Not a public SDK header. +#include "../../src/auth/oauth_internal.hpp" +#include "../support/resolve_gate.hpp" +#include "../support/socket_gate.hpp" + +#include +#include +#include +#include +#include +#include +#include + +#if BOOST_VERSION >= 107700 +#include +#include +#endif + +namespace asio = boost::asio; +namespace beast = boost::beast; +namespace http = beast::http; +using json = nlohmann::json; + +namespace { + +/// Loopback HTTP server that answers a scripted handler and records everything it received. +/// +/// `accepts()` is the socket ledger the security tests assert on: a control that refuses a target +/// must leave it at zero even though this server is listening and would happily accept. +class LoopbackServer final { + public: + using Handler = + std::function(const http::request&)>; + + explicit LoopbackServer(asio::io_context& io_ctx) + : acceptor_(io_ctx, {asio::ip::make_address("127.0.0.1"), 0}) {} + + [[nodiscard]] unsigned short port() const { return acceptor_.local_endpoint().port(); } + [[nodiscard]] std::string base_url() const { return "http://127.0.0.1:" + std::to_string(port()); } + [[nodiscard]] std::string origin() const { return base_url(); } + + void set_handler(Handler handler) { handler_ = std::move(handler); } + + [[nodiscard]] const std::vector& targets() const { return targets_; } + [[nodiscard]] const std::vector& bodies() const { return bodies_; } + [[nodiscard]] const std::vector& authorizations() const { return authorizations_; } + [[nodiscard]] int accepts() const { return accepts_; } + + /// Serve at most `request_budget` requests, then stop. close() aborts a pending accept so the + /// io_context always drains. + asio::awaitable serve(int request_budget) { + for (int index = 0; index < request_budget; ++index) { + boost::system::error_code accept_error; + auto socket = co_await acceptor_.async_accept( + asio::redirect_error(asio::use_awaitable, accept_error)); + if (accept_error) { + co_return; + } + ++accepts_; + + beast::tcp_stream stream(std::move(socket)); + beast::flat_buffer buffer; + http::request request; + boost::system::error_code read_error; + co_await http::async_read(stream, buffer, request, + asio::redirect_error(asio::use_awaitable, read_error)); + if (read_error) { + co_return; + } + + targets_.emplace_back(request.target()); + bodies_.push_back(request.body()); + authorizations_.emplace_back(request[http::field::authorization]); + + auto response = handler_(request); + response.version(request.version()); + response.prepare_payload(); + boost::system::error_code write_error; + co_await http::async_write(stream, response, + asio::redirect_error(asio::use_awaitable, write_error)); + + beast::error_code shutdown_error; + (void)stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + } + } + + void close() { + boost::system::error_code ignored; + (void)acceptor_.close(ignored); + } + + private: + asio::ip::tcp::acceptor acceptor_; + Handler handler_; + std::vector targets_; + std::vector bodies_; + std::vector authorizations_; + int accepts_{0}; +}; + +http::response json_response(const json& body) { + http::response response{http::status::ok, 11}; + response.set(http::field::content_type, "application/json"); + response.body() = body.dump(); + return response; +} + +http::response status_response(http::status status) { + http::response response{status, 11}; + response.body() = "{}"; + return response; +} + +json auth_server_metadata(const std::string& base, bool iss_supported) { + json metadata = {{"issuer", base}, + {"authorization_endpoint", base + "/authorize"}, + {"token_endpoint", base + "/token"}, + {"response_types_supported", json::array({"code"})}, + {"code_challenge_methods_supported", json::array({"S256"})}}; + if (iss_supported) { + metadata["authorization_response_iss_parameter_supported"] = true; + } + return metadata; +} + +json token_document() { + return {{"access_token", "granted-access-token"}, + {"token_type", "Bearer"}, + {"refresh_token", "granted-refresh-token"}, + {"expires_in", 3600}}; +} + +/// Extract one query parameter from an authorization URL. +std::string query_value(const std::string& url, const std::string& name) { + const auto response = mcp::auth::parse_authorization_response(url); + if (name == "state") { + return response.state.value_or(""); + } + const auto needle = name + "="; + auto position = url.find("?" + needle); + if (position == std::string::npos) { + position = url.find("&" + needle); + } + if (position == std::string::npos) { + return {}; + } + const auto start = position + needle.size() + 1; + const auto end = url.find('&', start); + return url.substr(start, end == std::string::npos ? std::string::npos : end - start); +} + +mcp::auth::MetadataFetchPolicy loopback_policy(const std::string& origin) { + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins.push_back(origin); + // The narrow opt-out the loopback fixture needs; never enabled implicitly. + policy.allow_plain_http_loopback = true; + return policy; +} + +/// Consent callback that echoes the recorded state and issuer, as a compliant server would. +mcp::auth::AuthorizationCallback echoing_callback(std::string* captured_url) { + return [captured_url](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + if (captured_url != nullptr) { + *captured_url = request.authorization_url; + } + mcp::auth::AuthorizationResponse response; + response.code = "test-authorization-code"; + response.state = request.state; + response.iss = request.issuer; + co_return response; + }; +} + +struct ManagerFixture { + std::shared_ptr store = + std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + std::string authorization_url; + int resolver_calls{0}; +}; + +} // namespace + +TEST(AuthAuthorizationManagerTest, CompletesChallengeDiscoveryAuthorizeAndTokenExchange) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/custom/prm.json") { + return json_response({{"resource", base + "/mcp"}, + {"authorization_servers", json::array({base})}, + {"scopes_supported", json::array({"mcp:read", "mcp:write"})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(3), asio::detached); + + ManagerFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.client_id = "test-client"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.policy = loopback_policy(server.origin()); + + bool authorized = false; + std::exception_ptr failure; + std::optional record; + std::string access_token; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.store, + fixture.config, + echoing_callback(&fixture.authorization_url)); + try { + authorized = co_await manager.try_handle_challenge( + R"(Bearer realm="mcp", resource_metadata=")" + base + R"(/custom/prm.json")"); + record = manager.last_authorization_request(); + access_token = manager.get_access_token(); + } catch (...) { + failure = std::current_exception(); + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_EQ(failure, nullptr); + EXPECT_TRUE(authorized); + ASSERT_TRUE(record.has_value()); + EXPECT_EQ(record->issuer, base); + EXPECT_TRUE(record->issuer_parameter_supported); + EXPECT_FALSE(record->state.empty()); + EXPECT_EQ(record->code_verifier.size(), 64U); + EXPECT_EQ(record->resource, base + "/mcp"); + EXPECT_EQ(record->scope, "mcp:read mcp:write"); + EXPECT_EQ(access_token, "granted-access-token"); + + // The challenge named the metadata location, so it is fetched directly and the well-known + // fallback is never probed. + ASSERT_EQ(server.targets().size(), 3U); + EXPECT_EQ(server.targets()[0], "/custom/prm.json"); + EXPECT_EQ(server.targets()[1], "/.well-known/oauth-authorization-server"); + EXPECT_EQ(server.targets()[2], "/token"); + + EXPECT_NE(fixture.authorization_url.find("code_challenge_method=S256"), std::string::npos); + EXPECT_NE(fixture.authorization_url.find("code_challenge="), std::string::npos); + EXPECT_NE(fixture.authorization_url.find("response_type=code"), std::string::npos); + EXPECT_EQ(query_value(fixture.authorization_url, "resource"), + mcp::auth::detail::url_encode(base + "/mcp")); + + const auto& token_body = server.bodies()[2]; + EXPECT_NE(token_body.find("grant_type=authorization_code"), std::string::npos); + EXPECT_NE(token_body.find("code=test-authorization-code"), std::string::npos); + EXPECT_NE(token_body.find("code_verifier="), std::string::npos); + EXPECT_NE(token_body.find("resource="), std::string::npos); +} + +TEST(AuthAuthorizationManagerTest, FallsBackToWellKnownOrderWhenTheChallengeNamesNoMetadata) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/.well-known/oauth-protected-resource/mcp") { + return status_response(http::status::not_found); + } + if (target == "/.well-known/oauth-protected-resource") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, false)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(4), asio::detached); + + ManagerFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.client_id = "test-client"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.policy = loopback_policy(server.origin()); + + bool authorized = false; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.store, + fixture.config, echoing_callback(nullptr)); + authorized = co_await manager.try_handle_challenge(R"(Bearer realm="mcp")"); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(authorized); + ASSERT_EQ(server.targets().size(), 4U); + // Path-based location is tried before the root one. + EXPECT_EQ(server.targets()[0], "/.well-known/oauth-protected-resource/mcp"); + EXPECT_EQ(server.targets()[1], "/.well-known/oauth-protected-resource"); +} + +TEST(AuthAuthorizationManagerTest, RejectsAnIssuerMismatchBeforeExchangingTheCode) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(2), asio::detached); + + ManagerFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.client_id = "test-client"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.policy = loopback_policy(server.origin()); + + std::string message; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + auto callback = [](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + mcp::auth::AuthorizationResponse response; + response.code = "test-authorization-code"; + response.state = request.state; + // Equivalent under RFC 3986 normalization, but the comparison is not canonicalizing. + response.iss = request.issuer + "/"; + co_return response; + }; + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.store, + fixture.config, callback); + try { + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm")"); + } catch (const std::exception& error) { + message = error.what(); + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_NE(message.find("iss did not match"), std::string::npos); + // The token endpoint was never reached. + EXPECT_EQ(server.targets().size(), 2U); + EXPECT_TRUE(fixture.store->load(base + "/mcp") == std::nullopt); +} + +TEST(AuthAuthorizationManagerTest, RefreshesTheStoredTokenAfterAuthorization) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + int token_requests = 0; + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + ++token_requests; + if (token_requests == 1) { + return json_response(token_document()); + } + return json_response({{"access_token", "renewed-access-token"}, + {"token_type", "Bearer"}, + {"expires_in", 3600}}); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(4), asio::detached); + + ManagerFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.client_id = "test-client"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.policy = loopback_policy(server.origin()); + + bool refreshed = false; + std::string token_after_refresh; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.store, + fixture.config, echoing_callback(nullptr)); + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm")"); + refreshed = co_await manager.try_refresh_token(); + token_after_refresh = manager.get_access_token(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(refreshed); + EXPECT_EQ(token_after_refresh, "renewed-access-token"); + EXPECT_EQ(token_requests, 2); + // The refresh grant reuses the refresh token issued with the original grant. + EXPECT_NE(server.bodies()[3].find("grant_type=refresh_token"), std::string::npos); + EXPECT_NE(server.bodies()[3].find("refresh_token=granted-refresh-token"), std::string::npos); +} + +TEST(AuthDiscoveryChallengeUrlTest, ChallengeSuppliedUrlSuppressesTheWellKnownFallback) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/custom/location.json") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + // Any well-known probe would be a priority-order violation, so report it as such. + return json_response({{"resource", "well-known-should-not-be-probed"}, + {"authorization_servers", json::array({base})}}); + }); + asio::co_spawn(io_ctx, server.serve(2), asio::detached); + + std::string resource; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(server.origin())); + mcp::auth::OAuthDiscoveryClient discovery(client); + const auto metadata = co_await discovery.discover_protected_resource( + base + "/mcp", std::optional(base + "/custom/location.json")); + resource = metadata.resource; + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_EQ(resource, base + "/mcp"); + ASSERT_EQ(server.targets().size(), 1U); + EXPECT_EQ(server.targets()[0], "/custom/location.json"); +} + +namespace { + +/// An Authenticator written before the challenge hook existed: it overrides only the two original +/// pure virtuals and must keep working unchanged. +class LegacyAuthenticator final : public mcp::auth::Authenticator { + public: + [[nodiscard]] std::string get_access_token() const override { return access_token_; } + + mcp::Task try_refresh_token() override { + ++refreshes_; + access_token_ = "legacy-refreshed-token"; + co_return true; + } + + [[nodiscard]] int refreshes() const { return refreshes_; } + + private: + std::string access_token_; + int refreshes_{0}; +}; + +} // namespace + +TEST(AuthChallengeReplayTest, LegacyAuthenticatorsStillFallThroughToRefresh) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + + server.set_handler([](const http::request& request) { + if (request[http::field::authorization].empty()) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer realm="mcp", resource_metadata="https://blocked.test/prm")"); + return challenge; + } + return status_response(http::status::accepted); + }); + asio::co_spawn(io_ctx, server.serve(2), asio::detached); + + auto authenticator = std::make_shared(); + const std::string wire = R"({"jsonrpc":"2.0","id":7,"method":"ping"})"; + std::exception_ptr failure; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + auto inner = std::make_shared(io_ctx.get_executor(), + server.base_url() + "/mcp"); + mcp::auth::OAuthClientTransport transport(inner, authenticator); + try { + co_await transport.write_message(wire); + } catch (...) { + failure = std::current_exception(); + } + transport.close(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_EQ(failure, nullptr); + // The default challenge hook reported that it handled nothing, so the legacy refresh ran. + EXPECT_EQ(authenticator->refreshes(), 1); + ASSERT_EQ(server.bodies().size(), 2U); + EXPECT_EQ(server.bodies()[1], wire); + EXPECT_EQ(server.authorizations()[1], "Bearer legacy-refreshed-token"); +} + +namespace { + +/// Drive one challenge against a manager whose resolver is instrumented, and report the refusal. +struct RefusalOutcome { + bool threw{false}; + mcp::auth::MetadataUrlDecision decision{mcp::auth::MetadataUrlDecision::allowed}; + int resolver_calls{0}; +}; + +RefusalOutcome refuse_challenge(const std::string& challenge, + const std::vector& extra_origins) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + + server.set_handler( + [](const http::request&) { return status_response(http::status::ok); }); + asio::co_spawn(io_ctx, server.serve(1), asio::detached); + + RefusalOutcome outcome; + auto store = std::make_shared(); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = server.base_url() + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + for (const auto& origin : extra_origins) { + config.policy.allowed_origins.push_back(origin); + } + config.host_resolver = [&outcome](const std::string&, + const std::string&) -> std::vector { + ++outcome.resolver_calls; + return {"127.0.0.1"}; + }; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + try { + (void)co_await manager.try_handle_challenge(challenge); + } catch (const mcp::auth::MetadataPolicyError& error) { + outcome.threw = true; + outcome.decision = error.decision(); + } catch (...) { + // Reported as a non-refusal by leaving `threw` false. + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + return outcome; +} + +} // namespace + +// The link-local metadata service is the canonical SSRF target. It must be refused before the host +// is ever resolved, and resolution is the only route to a socket. +TEST(AuthMetadataSsrfTest, RefusesTheLinkLocalMetadataServiceWithoutResolving) { + const auto outcome = + refuse_challenge(R"(Bearer resource_metadata="http://169.254.169.254/latest/meta-data")", + {"http://169.254.169.254"}); + EXPECT_TRUE(outcome.threw); + EXPECT_EQ(outcome.decision, mcp::auth::MetadataUrlDecision::scheme_not_allowed); + EXPECT_EQ(outcome.resolver_calls, 0); +} + +TEST(AuthMetadataSsrfTest, RefusesAPrivateRangeTargetWithoutResolving) { + const auto outcome = refuse_challenge(R"(Bearer resource_metadata="http://10.10.10.10/prm")", + {"http://10.10.10.10"}); + EXPECT_TRUE(outcome.threw); + EXPECT_EQ(outcome.decision, mcp::auth::MetadataUrlDecision::scheme_not_allowed); + EXPECT_EQ(outcome.resolver_calls, 0); +} + +TEST(AuthMetadataSsrfTest, RefusesPlainHttpToANonLoopbackHostWithoutResolving) { + const auto outcome = + refuse_challenge(R"(Bearer resource_metadata="http://metadata.example.test/prm")", + {"http://metadata.example.test"}); + EXPECT_TRUE(outcome.threw); + EXPECT_EQ(outcome.decision, mcp::auth::MetadataUrlDecision::scheme_not_allowed); + EXPECT_EQ(outcome.resolver_calls, 0); +} + +TEST(AuthMetadataSsrfTest, RefusesAnUnlistedOriginWithoutResolving) { + const auto outcome = + refuse_challenge(R"(Bearer resource_metadata="https://unlisted.example.test/prm")", {}); + EXPECT_TRUE(outcome.threw); + EXPECT_EQ(outcome.decision, mcp::auth::MetadataUrlDecision::origin_not_allowed); + EXPECT_EQ(outcome.resolver_calls, 0); +} + +// Address-level refusal. The server below is live and would accept instantly, and one of the two +// resolved answers points straight at it, so a non-zero accept count means a socket was opened +// despite the other answer being blocked. +TEST(AuthMetadataSsrfTest, RefusesTheWholeAnswerWhenOneResolvedAddressIsBlocked) { + asio::io_context io_ctx; + LoopbackServer canary(io_ctx); + + canary.set_handler([](const http::request&) { + return json_response({{"resource", "http://unused"}}); + }); + asio::co_spawn(io_ctx, canary.serve(1), asio::detached); + + int resolver_calls = 0; + bool refused = false; + auto decision = mcp::auth::MetadataUrlDecision::allowed; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + auto policy = loopback_policy("http://localhost:" + std::to_string(canary.port())); + client.set_metadata_policy(policy); + client.set_host_resolver( + [&](const std::string&, const std::string&) -> std::vector { + ++resolver_calls; + return {"127.0.0.1", "169.254.169.254"}; + }); + try { + (void)co_await client.get_json("http://localhost:" + std::to_string(canary.port()) + + "/prm"); + } catch (const mcp::auth::MetadataPolicyError& error) { + refused = true; + decision = error.decision(); + } + canary.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(refused); + EXPECT_EQ(decision, mcp::auth::MetadataUrlDecision::address_link_local); + EXPECT_EQ(resolver_calls, 1); + EXPECT_EQ(canary.accepts(), 0); +} + +// Resolve-then-pin: the addresses of one lookup are used for the connection and no second lookup +// happens, so a name that rebinds after the first answer cannot redirect the fetch. +TEST(AuthMetadataSsrfTest, PinsTheFirstResolutionSoRebindingCannotRedirectTheFetch) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + + server.set_handler([](const http::request&) { + return json_response( + {{"resource", "http://pinned"}, {"authorization_servers", json::array({"http://pinned"})}}); + }); + asio::co_spawn(io_ctx, server.serve(1), asio::detached); + + int resolver_calls = 0; + std::string resolved_resource; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + client.set_metadata_policy( + loopback_policy("http://localhost:" + std::to_string(server.port()))); + client.set_host_resolver( + [&](const std::string&, const std::string&) -> std::vector { + ++resolver_calls; + // A rebinding resolver: benign first, hostile on any later lookup. + if (resolver_calls == 1) { + return {"127.0.0.1"}; + } + return {"169.254.169.254"}; + }); + const auto document = + co_await client.get_json("http://localhost:" + std::to_string(server.port()) + "/prm"); + resolved_resource = document.at("resource").get(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_EQ(resolved_resource, "http://pinned"); + EXPECT_EQ(resolver_calls, 1); + EXPECT_EQ(server.accepts(), 1); +} + +TEST(AuthMetadataRedirectTest, RefusesARedirectOntoABlockedTarget) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + + server.set_handler([](const http::request&) { + http::response response{http::status::found, 11}; + response.set(http::field::location, "http://169.254.169.254/latest/meta-data"); + return response; + }); + asio::co_spawn(io_ctx, server.serve(2), asio::detached); + + bool refused = false; + auto decision = mcp::auth::MetadataUrlDecision::allowed; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + auto policy = loopback_policy(server.origin()); + policy.allowed_origins.emplace_back("http://169.254.169.254"); + client.set_metadata_policy(policy); + try { + (void)co_await client.get_json(server.base_url() + "/prm"); + } catch (const mcp::auth::MetadataPolicyError& error) { + refused = true; + decision = error.decision(); + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(refused); + EXPECT_EQ(decision, mcp::auth::MetadataUrlDecision::scheme_not_allowed); + // Only the original hop was made; the redirect target was refused rather than followed. + EXPECT_EQ(server.accepts(), 1); +} + +TEST(AuthMetadataRedirectTest, BoundsTheRedirectChain) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + + server.set_handler([](const http::request&) { + http::response response{http::status::found, 11}; + response.set(http::field::location, "/again"); + return response; + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + bool refused = false; + auto decision = mcp::auth::MetadataUrlDecision::allowed; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + auto policy = loopback_policy(server.origin()); + policy.max_redirects = 2; + client.set_metadata_policy(policy); + try { + (void)co_await client.get_json(server.base_url() + "/prm"); + } catch (const mcp::auth::MetadataPolicyError& error) { + refused = true; + decision = error.decision(); + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(refused); + EXPECT_EQ(decision, mcp::auth::MetadataUrlDecision::redirect_limit_exceeded); + // The initial request plus exactly two redirects. + EXPECT_EQ(server.accepts(), 3); +} + +TEST(AuthMetadataSizeCapTest, RejectsAnOversizedMetadataResponse) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + + server.set_handler([](const http::request&) { + http::response response{http::status::ok, 11}; + response.set(http::field::content_type, "application/json"); + response.body() = json{{"padding", std::string(64 * 1024, 'x')}}.dump(); + return response; + }); + asio::co_spawn(io_ctx, server.serve(1), asio::detached); + + bool refused = false; + auto decision = mcp::auth::MetadataUrlDecision::allowed; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + auto policy = loopback_policy(server.origin()); + policy.max_response_bytes = 4096; + client.set_metadata_policy(policy); + try { + (void)co_await client.get_json(server.base_url() + "/prm"); + } catch (const mcp::auth::MetadataPolicyError& error) { + refused = true; + decision = error.decision(); + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(refused); + EXPECT_EQ(decision, mcp::auth::MetadataUrlDecision::response_too_large); +} + +TEST(AuthMetadataSizeCapTest, AcceptsAResponseWithinTheCap) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + + server.set_handler([](const http::request&) { + return json_response({{"resource", "http://small"}}); + }); + asio::co_spawn(io_ctx, server.serve(1), asio::detached); + + std::string resource; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + auto policy = loopback_policy(server.origin()); + policy.max_response_bytes = 4096; + client.set_metadata_policy(policy); + const auto document = co_await client.get_json(server.base_url() + "/prm"); + resource = document.at("resource").get(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + EXPECT_EQ(resource, "http://small"); +} + +// End-to-end: an unauthenticated MCP request draws a challenge, the challenge drives a full +// authorization exchange, and the original request is replayed byte-for-byte with the new token. +TEST(AuthChallengeReplayTest, ReplaysTheExactRequestAfterChallengeDrivenAuthorization) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/mcp") { + if (request[http::field::authorization].empty()) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer realm="mcp", resource_metadata=")" + base + R"(/prm")"); + return challenge; + } + return status_response(http::status::accepted); + } + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + const std::string wire = R"({"jsonrpc":"2.0","id":1,"method":"ping"})"; + std::exception_ptr failure; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + auto inner = + std::make_shared(io_ctx.get_executor(), base + "/mcp"); + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, echoing_callback(nullptr)); + mcp::auth::OAuthClientTransport transport(inner, manager); + try { + co_await transport.write_message(wire); + } catch (...) { + failure = std::current_exception(); + } + transport.close(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_EQ(failure, nullptr); + ASSERT_EQ(server.targets().size(), 5U); + EXPECT_EQ(server.targets()[0], "/mcp"); + EXPECT_EQ(server.targets()[1], "/prm"); + EXPECT_EQ(server.targets()[2], "/.well-known/oauth-authorization-server"); + EXPECT_EQ(server.targets()[3], "/token"); + EXPECT_EQ(server.targets()[4], "/mcp"); + + // The replay is the same request, not a reconstruction, and it now carries the acquired token. + EXPECT_EQ(server.bodies()[4], server.bodies()[0]); + EXPECT_EQ(server.bodies()[4], wire); + EXPECT_TRUE(server.authorizations()[0].empty()); + EXPECT_EQ(server.authorizations()[4], "Bearer granted-access-token"); +} + +namespace { + +/// Consent callback that records the scope of every authorization request it is asked to run. +mcp::auth::AuthorizationCallback recording_callback(std::vector* scopes) { + return [scopes](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + scopes->push_back(request.scope.value_or("")); + mcp::auth::AuthorizationResponse response; + response.code = "test-authorization-code"; + response.state = request.state; + response.iss = request.issuer; + co_return response; + }; +} + +/// True when every space-separated scope in `needle` appears in `haystack`. +bool has_scopes(const std::string& haystack, const std::vector& needle) { + std::vector present; + std::string current; + for (const char character : haystack) { + if (character == ' ') { + if (!current.empty()) { + present.push_back(current); + } + current.clear(); + continue; + } + current.push_back(character); + } + if (!current.empty()) { + present.push_back(current); + } + for (const auto& wanted : needle) { + if (std::find(present.begin(), present.end(), wanted) == present.end()) { + return false; + } + } + return true; +} + +/// Accepts one connection and holds it open without ever reading or writing on it, modelling a +/// discovery or token endpoint that stalls -- the target `close()` must be able to abort without +/// waiting for it to time out on its own. +class StallingServer final { + public: + explicit StallingServer(asio::io_context& io_ctx) + : acceptor_(io_ctx, {asio::ip::make_address("127.0.0.1"), 0}) {} + + [[nodiscard]] unsigned short port() const { return acceptor_.local_endpoint().port(); } + [[nodiscard]] std::string base_url() const { return "http://127.0.0.1:" + std::to_string(port()); } + + /// Start accepting; the accepted socket is held as a member so the connection stays open (no + /// FIN, no RST) until this server is destroyed. + void accept_and_stall() { + acceptor_.async_accept([this](boost::system::error_code error, asio::ip::tcp::socket socket) { + if (!error) { + held_socket_ = std::move(socket); + accepted_.fetch_add(1); + } + }); + } + + /// Safe to poll from a thread other than the one running the io_context. + [[nodiscard]] int accepted() const { return accepted_.load(); } + + private: + asio::ip::tcp::acceptor acceptor_; + std::optional held_socket_; + std::atomic accepted_{0}; +}; + +} // namespace + +TEST(AuthScopeStepUpTest, UnionsTheGrantedScopeWithAForbiddenChallengeAndReplaysTheRequest) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + int mcp_calls = 0; + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/mcp") { + ++mcp_calls; + if (mcp_calls == 1) { + // Unauthenticated: the challenge names a scope that is disjoint from the PRM's + // `scopes_supported`, and the challenge must still win. + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer scope="mcp:basic", resource_metadata=")" + base + R"(/prm")"); + return challenge; + } + if (mcp_calls == 2) { + // Step-up: the token is valid, but this operation needs a scope the grant lacks. + http::response challenge{http::status::forbidden, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer error="insufficient_scope", scope="mcp:write", )" + R"(resource_metadata=")" + + base + R"(/prm")"); + return challenge; + } + return status_response(http::status::accepted); + } + if (target == "/prm") { + return json_response({{"resource", base + "/mcp"}, + {"authorization_servers", json::array({base})}, + {"scopes_supported", json::array({"prm:only"})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(20), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + std::vector requested_scopes; + const std::string wire = R"({"jsonrpc":"2.0","id":1,"method":"tools/call"})"; + std::exception_ptr failure; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + auto inner = + std::make_shared(io_ctx.get_executor(), base + "/mcp"); + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + mcp::auth::OAuthClientTransport transport(inner, manager); + try { + co_await transport.write_message(wire); + } catch (...) { + failure = std::current_exception(); + } + transport.close(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_EQ(failure, nullptr); + ASSERT_EQ(requested_scopes.size(), 2U); + + // The challenge scope is authoritative even though it shares nothing with `scopes_supported`: + // no set relationship between the two may be assumed in either direction. + EXPECT_EQ(requested_scopes[0], "mcp:basic"); + EXPECT_FALSE(has_scopes(requested_scopes[0], {"prm:only"})); + + // Step-up re-authorizes on the union, so the scope already granted survives the escalation. + EXPECT_TRUE(has_scopes(requested_scopes[1], {"mcp:basic", "mcp:write"})); + EXPECT_FALSE(has_scopes(requested_scopes[1], {"prm:only"})); + + // The replay is the same request bytes, carrying the newly acquired token. + EXPECT_EQ(mcp_calls, 3); + ASSERT_FALSE(server.bodies().empty()); + EXPECT_EQ(server.bodies().back(), wire); + EXPECT_EQ(server.authorizations().back(), "Bearer granted-access-token"); +} + +TEST(AuthScopeStepUpTest, StopsAfterThreeAuthorizationChallengesForOneRequest) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/mcp") { + // A scope escalation that will never succeed, exactly as the retry-limit fixture does. + if (request[http::field::authorization].empty()) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer scope="mcp:admin", resource_metadata=")" + base + R"(/prm")"); + return challenge; + } + http::response challenge{http::status::forbidden, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer error="insufficient_scope", scope="mcp:admin", )" + R"(resource_metadata=")" + + base + R"(/prm")"); + return challenge; + } + if (target == "/prm") { + return json_response({{"resource", base + "/mcp"}, + {"authorization_servers", json::array({base})}, + {"scopes_supported", json::array({"mcp:admin"})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(40), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + std::vector requested_scopes; + std::exception_ptr failure; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + auto inner = + std::make_shared(io_ctx.get_executor(), base + "/mcp"); + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + mcp::auth::OAuthClientTransport transport(inner, manager); + try { + co_await transport.write_message(R"({"jsonrpc":"2.0","id":1,"method":"tools/call"})"); + } catch (...) { + failure = std::current_exception(); + } + transport.close(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + // The refusal is surfaced rather than retried forever, and the cap is three per request. + EXPECT_NE(failure, nullptr); + EXPECT_EQ(requested_scopes.size(), 3U); +} + +TEST(AuthScopeStepUpTest, CoalescesConcurrentChallengesIntoASingleAuthorizationFlow) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(10), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + std::vector requested_scopes; + const std::string header = R"(Bearer scope="mcp:basic", resource_metadata=")" + base + R"(/prm")"; + + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + + int completed = 0; + int succeeded = 0; + const auto challenger = [&]() -> mcp::Task { + const auto authorized = co_await manager->try_handle_challenge(header); + succeeded += authorized ? 1 : 0; + if (++completed == 2) { + server.close(); + } + }; + asio::co_spawn(io_ctx, challenger(), asio::detached); + asio::co_spawn(io_ctx, challenger(), asio::detached); + + io_ctx.run(); + + // Both callers are authorized, but only one of them ran a flow. + EXPECT_EQ(completed, 2); + EXPECT_EQ(succeeded, 2); + EXPECT_EQ(requested_scopes.size(), 1U); + EXPECT_EQ(store->load(config.server_url).has_value(), true); +} + +TEST(AuthTransportCloseTest, CloseDuringDiscoveryAbortsTheStalledExchange) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto stalling_base = stalling.base_url(); + + // The resource server answers instantly; only the metadata target it names stalls, so it is + // discovery -- not the initial request -- that close() must abort. + LoopbackServer resource_server(io_ctx); + const auto resource_base = resource_server.base_url(); + resource_server.set_handler([&](const http::request&) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer resource_metadata=")" + stalling_base + R"(/prm")"); + return challenge; + }); + asio::co_spawn(io_ctx, resource_server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = resource_base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {resource_base, stalling_base}; + policy.allow_plain_http_loopback = true; + config.policy = policy; + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + auto inner = + std::make_shared(io_ctx.get_executor(), resource_base + "/mcp"); + auto transport = std::make_shared(inner, manager); + + bool completed = false; + bool timed_out = false; + std::exception_ptr failure; + + // Declared before the work is spawned so the coroutine can cancel it on completion instead of + // io_ctx.run() always waiting out the full watchdog window. + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(10)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && !completed) { + timed_out = true; + io_ctx.stop(); + } + }); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"tools/call"})"); + } catch (...) { + failure = std::current_exception(); + } + completed = true; + watchdog.cancel(); + }, + asio::detached); + + // Close once the write has had time to reach the stalled discovery fetch. + asio::steady_timer closer(io_ctx); + closer.expires_after(std::chrono::milliseconds(200)); + closer.async_wait([&](boost::system::error_code) { + transport->close(); + resource_server.close(); + }); + + io_ctx.run(); + + ASSERT_FALSE(timed_out) << "watchdog: close() did not unblock the stalled discovery exchange"; + EXPECT_TRUE(completed); + EXPECT_NE(failure, nullptr); + EXPECT_GE(stalling.accepted(), 1); +} + +TEST(AuthTransportCloseTest, CloseDuringTokenExchangeAbortsTheStalledExchange) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto stalling_base = stalling.base_url(); + + // Discovery and the consent redirect both succeed normally; only the token endpoint -- pointed + // at the stalling server -- never answers, so close() must abort the token exchange itself. + LoopbackServer resource_server(io_ctx); + const auto resource_base = resource_server.base_url(); + resource_server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/mcp") { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer resource_metadata=")" + resource_base + R"(/prm")"); + return challenge; + } + if (target == "/prm") { + return json_response({{"resource", resource_base + "/mcp"}, + {"authorization_servers", json::array({resource_base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + json metadata = auth_server_metadata(resource_base, true); + metadata["token_endpoint"] = stalling_base + "/token"; + return json_response(metadata); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, resource_server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = resource_base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {resource_base, stalling_base}; + policy.allow_plain_http_loopback = true; + config.policy = policy; + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + auto inner = + std::make_shared(io_ctx.get_executor(), resource_base + "/mcp"); + auto transport = std::make_shared(inner, manager); + + bool completed = false; + bool timed_out = false; + std::exception_ptr failure; + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(10)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && !completed) { + timed_out = true; + io_ctx.stop(); + } + }); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"tools/call"})"); + } catch (...) { + failure = std::current_exception(); + } + completed = true; + watchdog.cancel(); + }, + asio::detached); + + asio::steady_timer closer(io_ctx); + closer.expires_after(std::chrono::milliseconds(200)); + closer.async_wait([&](boost::system::error_code) { + transport->close(); + resource_server.close(); + }); + + io_ctx.run(); + + ASSERT_FALSE(timed_out) << "watchdog: close() did not unblock the stalled token exchange"; + EXPECT_TRUE(completed); + EXPECT_NE(failure, nullptr); + EXPECT_GE(stalling.accepted(), 1); +} + +// A follower must report the outcome of the flight it joined, with the leader's own failure, not the +// result of a later flight. The sequencing is arranged rather than raced: +// +// 1. The leader of flight one parks inside the application's authorization callback. +// 2. A follower joins flight one and parks on its timer. +// 3. That follower's own strand is then blocked, so its wake-up cannot be delivered. +// 4. The leader of flight one is released and fails with a distinctive message. +// 5. A second challenge runs to completion and SUCCEEDS. +// 6. Only then is the follower's strand released and its result read. +// +// At step 6 the follower must fail with flight one's message, not report flight two's success. +TEST(AuthTransportCloseTest, AFollowerReportsItsOwnFlightsFailureNotALaterFlightsSuccess) { + constexpr int io_thread_count = 4; + const std::string leader_failure_text = "flight one refused by the application callback"; + + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm.json") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(6), asio::detached); + + // The first leader runs here, and so does the gate it parks on, so the gate is only ever + // touched from one strand. Expiring a timer from a thread that is not the one waiting on it is + // the very defect the single-flight timer had; the test must not reintroduce it. + auto leader_strand = asio::make_strand(io_ctx); + auto follower_strand = asio::make_strand(io_ctx); + asio::steady_timer leader_gate(leader_strand, asio::steady_timer::time_point::max()); + + std::atomic callback_calls{0}; + std::promise leader_parked_signal; + auto leader_parked = leader_parked_signal.get_future(); + + auto callback = [&](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + if (callback_calls.fetch_add(1) == 0) { + leader_parked_signal.set_value(); + boost::system::error_code ignored; + co_await leader_gate.async_wait(asio::redirect_error(asio::use_awaitable, ignored)); + throw std::runtime_error(leader_failure_text); + } + // The second challenge is a normal, successful authorization. + mcp::auth::AuthorizationResponse response; + response.code = "test-authorization-code"; + response.state = request.state; + response.iss = request.issuer; + co_return response; + }; + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto manager = std::make_shared(io_ctx.get_executor(), store, + config, callback); + const std::string header = R"(Bearer realm="mcp", resource_metadata=")" + base + R"(/prm.json")"; + + std::promise first_leader_signal; + auto first_leader = first_leader_signal.get_future(); + std::promise second_leader_signal; + auto second_leader = second_leader_signal.get_future(); + std::promise> follower_signal; + auto follower_result = follower_signal.get_future(); + + std::promise blocker_running_signal; + auto blocker_running = blocker_running_signal.get_future(); + std::promise blocker_release_signal; + auto blocker_release = blocker_release_signal.get_future(); + + std::atomic timed_out{false}; + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(30)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error) { + timed_out.store(true); + io_ctx.stop(); + } + }); + + std::vector runners; + runners.reserve(io_thread_count); + for (int index = 0; index < io_thread_count; ++index) { + runners.emplace_back([&io_ctx]() { io_ctx.run(); }); + } + + asio::co_spawn( + leader_strand, + [&]() -> mcp::Task { + std::string message; + try { + (void)co_await manager->try_handle_challenge(header); + message = ""; + } catch (const std::exception& error) { + message = error.what(); + } + first_leader_signal.set_value(message); + }, + asio::detached); + + ASSERT_EQ(leader_parked.wait_for(std::chrono::seconds(10)), std::future_status::ready) + << "the first leader never reached the authorization callback"; + + // Spawned before the blocker is posted to the same strand, so the follower has already joined + // the flight and suspended by the time the blocker takes the strand over. Joining happens + // synchronously inside try_handle_challenge(), before the coroutine's first suspension. + asio::co_spawn( + follower_strand, + [&]() -> mcp::Task { + bool authorized = false; + std::string message; + try { + authorized = co_await manager->try_handle_challenge(header); + } catch (const std::exception& error) { + message = error.what(); + } + follower_signal.set_value({authorized, message}); + }, + asio::detached); + + // Holds the follower's strand so its wake-up stays queued while the second flight runs and + // finishes. This is what makes "a follower that wakes late" deterministic instead of a race. + asio::post(follower_strand, [&]() { + blocker_running_signal.set_value(); + blocker_release.wait(); + }); + ASSERT_EQ(blocker_running.wait_for(std::chrono::seconds(10)), std::future_status::ready) + << "the follower's strand was never taken over, so nothing was held back"; + + // Release the first leader, which now fails. + asio::post(leader_strand, + [&leader_gate]() { leader_gate.expires_at(asio::steady_timer::time_point::min()); }); + ASSERT_EQ(first_leader.wait_for(std::chrono::seconds(10)), std::future_status::ready); + const auto first_message = first_leader.get(); + ASSERT_NE(first_message.find(leader_failure_text), std::string::npos) + << "the first leader did not fail the way this test needs it to: " << first_message; + + // A second, successful flight: the result a manager-wide field would hand the follower. + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + bool authorized = false; + try { + authorized = co_await manager->try_handle_challenge(header); + } catch (...) { + authorized = false; + } + second_leader_signal.set_value(authorized); + }, + asio::detached); + ASSERT_EQ(second_leader.wait_for(std::chrono::seconds(10)), std::future_status::ready); + ASSERT_TRUE(second_leader.get()) + << "the second flight had to succeed for this test to mean anything"; + + // Only now may the follower wake. + blocker_release_signal.set_value(); + const auto follower_status = follower_result.wait_for(std::chrono::seconds(10)); + + io_ctx.stop(); + for (auto& runner : runners) { + runner.join(); + } + + ASSERT_FALSE(timed_out.load()) << "watchdog fired"; + ASSERT_EQ(follower_status, std::future_status::ready) << "the follower never woke"; + const auto [follower_authorized, follower_message] = follower_result.get(); + + EXPECT_FALSE(follower_authorized) + << "the follower reported the LATER flight's success for a flight that failed"; + EXPECT_NE(follower_message.find(leader_failure_text), std::string::npos) + << "the follower did not surface its own leader's reason, it got: " << follower_message; +} + +TEST(AuthTransportCloseTest, CloseWhileAFollowerIsParkedOnTheSingleFlightTimerWakesItWithAnError) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto stalling_base = stalling.base_url(); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = stalling_base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {stalling_base}; + policy.allow_plain_http_loopback = true; + config.policy = policy; + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + + // Discovery itself is the stalled step: the leader parks inside it, and the follower parks on + // the single-flight timer behind the leader. + const std::string header = R"(Bearer resource_metadata=")" + stalling_base + R"(/prm")"; + + bool leader_done = false; + bool follower_done = false; + bool timed_out = false; + bool leader_authorized = false; + bool follower_authorized = false; + std::exception_ptr leader_failure; + std::exception_ptr follower_failure; + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(10)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && !(leader_done && follower_done)) { + timed_out = true; + io_ctx.stop(); + } + }); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + leader_authorized = co_await manager->try_handle_challenge(header); + } catch (...) { + leader_failure = std::current_exception(); + } + leader_done = true; + if (follower_done) { + watchdog.cancel(); + } + }, + asio::detached); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + follower_authorized = co_await manager->try_handle_challenge(header); + } catch (...) { + follower_failure = std::current_exception(); + } + follower_done = true; + if (leader_done) { + watchdog.cancel(); + } + }, + asio::detached); + + asio::steady_timer closer(io_ctx); + closer.expires_after(std::chrono::milliseconds(200)); + closer.async_wait([&](boost::system::error_code) { manager->close(); }); + + io_ctx.run(); + + ASSERT_FALSE(timed_out) << "watchdog: close() did not wake the parked follower"; + EXPECT_TRUE(leader_done); + EXPECT_TRUE(follower_done); + EXPECT_FALSE(leader_authorized); + EXPECT_FALSE(follower_authorized); + // The follower gets a clear error, not the leader's misleading "not authorized" outcome. + EXPECT_NE(follower_failure, nullptr); + EXPECT_NE(leader_failure, nullptr); +} + +// Stalls the leader where abort_pending() cannot reach -- the application's consent callback -- and +// races close() against a follower joining the same flight from a second OS thread, so the follower's +// join and its flight->async_wait() registration are not guaranteed to have both completed before +// close() runs. A bare flight->cancel() is a no-op against a wait that has not started, which would +// hang; expires_at(time_point::min()) makes a wait registered after the call complete immediately. +TEST(AuthTransportCloseTest, + CloseWhileALeaderIsParkedInTheApplicationConsentCallbackWakesAFollowerWithAnError) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + + // Never returns, modelling an application consent prompt nobody has answered yet. Signals + // `leader_parked_signal` right before parking, so the main thread knows `flight` already exists + // (it is created earlier still, synchronously, before discovery even starts) and the leader has + // reached the one phase this manager cannot itself abort. The leader is deliberately left parked + // here, leaked into the stopped io_context, for the rest of the test: only the follower's release + // is under test, so nothing below joins or waits on the leader's own coroutine. + std::promise leader_parked_signal; + auto leader_parked = leader_parked_signal.get_future(); + auto callback = + [&io_ctx, &leader_parked_signal]( + const mcp::auth::AuthorizationRequest&) -> mcp::Task { + leader_parked_signal.set_value(); + asio::steady_timer never(io_ctx, asio::steady_timer::time_point::max()); + boost::system::error_code ignored; + co_await never.async_wait(asio::redirect_error(asio::use_awaitable, ignored)); + co_return mcp::auth::AuthorizationResponse{}; + }; + + auto manager = std::make_shared(io_ctx.get_executor(), store, + config, callback); + const std::string header = R"(Bearer resource_metadata=")" + base + R"(/prm")"; + + // A single follower racing a single close() call almost never lands in the gap between joining + // the flight and registering its wait: starting the runner thread, having it work through + // discovery and reach the callback, and then waking it again for one posted follower all take far + // longer than close()'s own few instructions. A burst of many followers, posted individually and + // racing the same close() call from a second thread with no synchronization, gives the same + // narrow window many independent chances to be hit in one test run instead of one. + constexpr int follower_count = 200; + std::vector follower_authorized(follower_count, false); + std::vector follower_failure(follower_count); + int followers_done = 0; + bool timed_out = false; + std::exception_ptr leader_failure; + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(10)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && followers_done < follower_count) { + timed_out = true; + } + io_ctx.stop(); + }); + + // asio::detached swallows an uncaught exception silently, which would otherwise leave + // `leader_parked_signal` unfulfilled forever with nothing left to explain why; captured here + // purely as a diagnostic in case discovery itself fails. + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await manager->try_handle_challenge(header); + } catch (...) { + leader_failure = std::current_exception(); + } + }, + asio::detached); + + std::thread runner([&io_ctx]() { io_ctx.run(); }); + // Bounded even though discovery against this instant loopback server should resolve in well + // under a millisecond: nothing here may block the main thread indefinitely, since a stall here + // would sit outside the io_context's own watchdog entirely. + const auto parked_status = leader_parked.wait_for(std::chrono::seconds(5)); + + if (parked_status == std::future_status::ready) { + for (int index = 0; index < follower_count; ++index) { + asio::post(io_ctx, [&, index]() { + asio::co_spawn( + io_ctx, + [&, index]() -> mcp::Task { + try { + follower_authorized[index] = co_await manager->try_handle_challenge(header); + } catch (...) { + follower_failure[index] = std::current_exception(); + } + if (++followers_done == follower_count) { + io_ctx.stop(); + } + }, + asio::detached); + }); + } + // Deliberately no synchronization beyond what the manager itself provides: this call races + // the followers' posts above from a second thread that is not running the io_context at all, + // which is exactly how close() is used in practice (an application thread tearing down a + // transport while the io_context spins elsewhere). + manager->close(); + } + + runner.join(); + + ASSERT_EQ(parked_status, std::future_status::ready) + << "leader never reached the consent callback (discovery failed? " + << (leader_failure ? "yes, see leader_failure" : "no exception captured"); + ASSERT_FALSE(timed_out) << "watchdog: close() did not wake every follower (" << followers_done + << "/" << follower_count << " woke up)"; + EXPECT_EQ(followers_done, follower_count); + for (int index = 0; index < follower_count; ++index) { + EXPECT_FALSE(follower_authorized[index]) << "follower " << index; + EXPECT_NE(follower_failure[index], nullptr) << "follower " << index; + } +} + +// The single-flight timer's two touch points -- expire_flight()'s expires_at() and a follower's +// async_wait() in await_in_flight() -- run synchronously on whatever thread calls them, and a +// steady_timer is not safe for concurrent use. Two things are needed to reach the window: +// +// * The io_context runs on several threads, so the two calls can execute at the same instant. +// * Followers keep arriving while close() lands: close() sets `closed` first, so a follower +// spawned afterwards throws in handle_challenge() without reaching the timer. Only one that +// read `flight` before `closed` was set and calls async_wait() after expire_flight() ran is in +// the gap. +// +// Unserialised, ThreadSanitizer reports the race; without a sanitizer a follower whose wait was +// enqueued at time_point::max() is never woken, and the watchdog fires. +TEST(AuthTransportCloseTest, CloseWakesEveryFollowerWhileTheyKeepArrivingOnAMultiThreadedIoContext) { + constexpr int io_thread_count = 4; + // A strand serialises its own followers but not the followers on the other strands, so several + // strands give the narrow window many independent chances per run. + constexpr int follower_strand_count = 4; + constexpr int follower_count = 3000; + constexpr auto close_delay = std::chrono::microseconds(300); + + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + + // Parks forever, modelling an application consent prompt nobody has answered, so the flight + // stays open for followers to coalesce onto. Deliberately leaked into the stopped io_context + // exactly as the single-threaded consent test leaks it: releasing a parked leader is not what + // this test covers. + std::promise leader_parked_signal; + auto leader_parked = leader_parked_signal.get_future(); + auto callback = + [&io_ctx, &leader_parked_signal]( + const mcp::auth::AuthorizationRequest&) -> mcp::Task { + leader_parked_signal.set_value(); + asio::steady_timer never(io_ctx, asio::steady_timer::time_point::max()); + boost::system::error_code ignored; + co_await never.async_wait(asio::redirect_error(asio::use_awaitable, ignored)); + co_return mcp::auth::AuthorizationResponse{}; + }; + + auto manager = std::make_shared(io_ctx.get_executor(), store, + config, callback); + const std::string header = R"(Bearer resource_metadata=")" + base + R"(/prm")"; + + // Counters rather than per-index vectors: these are written from four io threads at once, and + // std::vector packs its elements into shared words, which would be a race in the test + // itself rather than in the code under test. + std::atomic followers_done{0}; + std::atomic followers_failed{0}; + std::atomic followers_authorized{0}; + std::atomic resumed_off_own_strand{0}; + std::atomic timed_out{false}; + std::exception_ptr leader_failure; + + // Each follower runs on one of these, standing in for the strand a real caller is on: Client + // spawns its write onto its own strand and SerializedTransportWriter builds another to + // serialise writes. Declared out here so the follower coroutines can still name their own + // strand after the block below has ended. + std::vector> strands; + strands.reserve(follower_strand_count); + for (int index = 0; index < follower_strand_count; ++index) { + strands.push_back(asio::make_strand(io_ctx)); + } + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(30)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && followers_done.load() < follower_count) { + timed_out.store(true); + } + io_ctx.stop(); + }); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await manager->try_handle_challenge(header); + } catch (...) { + leader_failure = std::current_exception(); + } + }, + asio::detached); + + std::vector runners; + runners.reserve(io_thread_count); + for (int index = 0; index < io_thread_count; ++index) { + runners.emplace_back([&io_ctx]() { io_ctx.run(); }); + } + + const auto parked_status = leader_parked.wait_for(std::chrono::seconds(10)); + + std::thread feeder; + if (parked_status == std::future_status::ready) { + feeder = std::thread([&]() { + for (int index = 0; index < follower_count; ++index) { + const int strand_index = index % follower_strand_count; + asio::co_spawn( + strands[strand_index], + [&, strand_index]() -> mcp::Task { + try { + if (co_await manager->try_handle_challenge(header)) { + followers_authorized.fetch_add(1); + } + } catch (...) { + followers_failed.fetch_add(1); + } + // await_in_flight() initiates its wait on the manager's flight strand, so + // the continuation has to be handed back to the follower's own executor; + // AFollowerReleasedFromTheSingleFlightTimerResumesOnItsOwnStrand below + // covers this unconditionally. + if (!strands[strand_index].running_in_this_thread()) { + resumed_off_own_strand.fetch_add(1); + } + if (followers_done.fetch_add(1) + 1 == follower_count) { + io_ctx.stop(); + } + }, + asio::detached); + } + }); + + // No synchronization beyond what the manager itself provides, from a thread that is not + // running the io_context: this is how an application tears a transport down. The delay + // decides where in the still-draining follower stream close() lands. + std::this_thread::sleep_for(close_delay); + manager->close(); + feeder.join(); + } + + for (auto& runner : runners) { + runner.join(); + } + + ASSERT_EQ(parked_status, std::future_status::ready) + << "leader never reached the consent callback (discovery failed? " + << (leader_failure ? "yes, see leader_failure" : "no exception captured") << ")"; + ASSERT_FALSE(timed_out.load()) + << "watchdog: a follower was left parked on the single-flight timer (" << followers_done.load() + << "/" << follower_count << " woke up)"; + EXPECT_EQ(followers_done.load(), follower_count); + // Every follower either joined the flight and was released by close(), or arrived after + // `closed` was set and threw straight away. Neither outcome authorizes anything: the leader + // never got past the consent callback. + EXPECT_EQ(followers_authorized.load(), 0); + EXPECT_EQ(followers_failed.load(), follower_count); + EXPECT_EQ(resumed_off_own_strand.load(), 0) + << "a follower resumed off the strand it was spawned on: await_in_flight() left the caller " + "on the manager's flight strand instead of handing it back"; +} + +// await_in_flight() initiates its wait on the manager's flight strand, which takes the follower off +// the executor it was spawned on. Everything after that wait has to come back: Client and +// SerializedTransportWriter each serialise writes on a strand of their own, and a continuation left +// on the flight strand would bypass both. Exactly one follower is used, proven to have parked before +// close() runs, so the executor assertion is unconditional. +TEST(AuthTransportCloseTest, AFollowerReleasedFromTheSingleFlightTimerResumesOnItsOwnStrand) { + constexpr int io_thread_count = 4; + + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + + std::promise leader_parked_signal; + auto leader_parked = leader_parked_signal.get_future(); + auto callback = + [&io_ctx, &leader_parked_signal]( + const mcp::auth::AuthorizationRequest&) -> mcp::Task { + leader_parked_signal.set_value(); + asio::steady_timer never(io_ctx, asio::steady_timer::time_point::max()); + boost::system::error_code ignored; + co_await never.async_wait(asio::redirect_error(asio::use_awaitable, ignored)); + co_return mcp::auth::AuthorizationResponse{}; + }; + + auto manager = std::make_shared(io_ctx.get_executor(), store, + config, callback); + const std::string header = R"(Bearer resource_metadata=")" + base + R"(/prm")"; + + // The follower's own strand, standing in for a real caller's. + auto follower_strand = asio::make_strand(io_ctx); + + std::atomic follower_done{false}; + std::atomic follower_threw{false}; + std::atomic follower_on_own_strand{false}; + std::atomic timed_out{false}; + std::exception_ptr leader_failure; + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(30)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && !follower_done.load()) { + timed_out.store(true); + } + io_ctx.stop(); + }); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await manager->try_handle_challenge(header); + } catch (...) { + leader_failure = std::current_exception(); + } + }, + asio::detached); + + std::vector runners; + runners.reserve(io_thread_count); + for (int index = 0; index < io_thread_count; ++index) { + runners.emplace_back([&io_ctx]() { io_ctx.run(); }); + } + + const auto parked_status = leader_parked.wait_for(std::chrono::seconds(10)); + + bool follower_was_parked = false; + if (parked_status == std::future_status::ready) { + asio::co_spawn( + follower_strand, + [&]() -> mcp::Task { + try { + (void)co_await manager->try_handle_challenge(header); + } catch (...) { + follower_threw.store(true); + } + follower_on_own_strand.store(follower_strand.running_in_this_thread()); + follower_done.store(true); + io_ctx.stop(); + }, + asio::detached); + + // Long enough for the follower to join the flight and register its wait. What makes this + // deterministic is not the sleep but the check after it: if the follower had returned + // without parking, it would already be done. + std::this_thread::sleep_for(std::chrono::milliseconds(250)); + follower_was_parked = !follower_done.load(); + + manager->close(); + } + + for (auto& runner : runners) { + runner.join(); + } + + ASSERT_EQ(parked_status, std::future_status::ready) + << "leader never reached the consent callback (discovery failed? " + << (leader_failure ? "yes, see leader_failure" : "no exception captured") << ")"; + ASSERT_TRUE(follower_was_parked) + << "the follower finished before close() ran, so it never parked on the single-flight " + "timer and this test proves nothing"; + ASSERT_FALSE(timed_out.load()) << "watchdog: close() never woke the parked follower"; + EXPECT_TRUE(follower_threw.load()) << "a follower released by close() must report an error"; + EXPECT_TRUE(follower_on_own_strand.load()) + << "the follower resumed off the strand it was spawned on: await_in_flight() left it on " + "the manager's flight strand instead of handing it back"; +} + +// flight->expires_at()/cancel() in close() and flight->async_wait() in await_in_flight() touch +// the same non-thread-safe timer object; close() reaches it synchronously from whatever thread the +// application calls OAuthClientTransport::close() from, which is not necessarily the thread running +// the io_context. Runs the io_context on its own thread and calls close() from the main thread with +// no synchronization beyond the manager's own, while a flow is genuinely in flight. Boost.Asio's +// timer and socket types are not safe under concurrent access from two threads, so this is the +// shape a sanitizer build targets. +TEST(AuthTransportCloseTest, CloseFromAnotherThreadWhileAuthorizationIsInFlightTerminatesCleanly) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto stalling_base = stalling.base_url(); + + LoopbackServer resource_server(io_ctx); + const auto resource_base = resource_server.base_url(); + resource_server.set_handler([&](const http::request&) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer resource_metadata=")" + stalling_base + R"(/prm")"); + return challenge; + }); + asio::co_spawn(io_ctx, resource_server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = resource_base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {resource_base, stalling_base}; + policy.allow_plain_http_loopback = true; + config.policy = policy; + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + auto inner = + std::make_shared(io_ctx.get_executor(), resource_base + "/mcp"); + auto transport = std::make_shared(inner, manager); + + bool completed = false; + bool timed_out = false; + std::exception_ptr failure; + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(10)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && !completed) { + timed_out = true; + } + io_ctx.stop(); + }); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"tools/call"})"); + } catch (...) { + failure = std::current_exception(); + } + completed = true; + io_ctx.stop(); + }, + asio::detached); + + std::thread runner([&io_ctx]() { io_ctx.run(); }); + + // Wait until the discovery fetch has connected to the stalling server (a fixed head start is + // not enough under sanitizer slowdown), then close from a thread that never runs the io_context. + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(5); + while (stalling.accepted() < 1 && std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + const bool reached_stall = stalling.accepted() >= 1; + transport->close(); + resource_server.close(); + + runner.join(); + + ASSERT_TRUE(reached_stall) + << "the flow never reached the stalling discovery server, so close() raced nothing"; + ASSERT_FALSE(timed_out) << "watchdog: close() from another thread did not unblock the flow"; + EXPECT_TRUE(completed); + EXPECT_NE(failure, nullptr); + EXPECT_GE(stalling.accepted(), 1); +} + +// Between the per-hop latch check and the socket opening inside async_connect() an exchange has +// nothing live for abort_pending() to close, so connect() re-reads the latch once the addresses are +// known. The resolver hook runs synchronously at exactly that point. The fetch reports +// operation_aborted either way, so the oracle is the listener's accept queue: a probe connected after +// the fetch has failed must be the first connection the listener hands out. +TEST(AuthTransportCloseTest, AbortBetweenResolutionAndConnectOpensNoConnection) { + asio::io_context io_ctx; + asio::ip::tcp::acceptor acceptor(io_ctx, {asio::ip::make_address("127.0.0.1"), 0}); + const auto listening = acceptor.local_endpoint(); + const auto base = "http://127.0.0.1:" + std::to_string(listening.port()); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(base)); + + int resolver_calls = 0; + auto* raw_client = client.get(); + client->set_host_resolver([&resolver_calls, raw_client](const std::string&, const std::string&) { + ++resolver_calls; + raw_client->abort_pending(); + return std::vector{"127.0.0.1"}; + }); + + bool completed = false; + bool timed_out = false; + boost::system::error_code fetch_error; + asio::ip::tcp::endpoint probe_endpoint; + asio::ip::tcp::endpoint first_accepted_endpoint; + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(10)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && !completed) { + timed_out = true; + io_ctx.stop(); + } + }); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await client->get_json(base + "/prm"); + } catch (const boost::system::system_error& error) { + fetch_error = error.code(); + } + + asio::ip::tcp::socket probe(io_ctx); + co_await probe.async_connect(listening, asio::use_awaitable); + probe_endpoint = probe.local_endpoint(); + + asio::ip::tcp::socket accepted(io_ctx); + co_await acceptor.async_accept(accepted, first_accepted_endpoint, asio::use_awaitable); + + completed = true; + watchdog.cancel(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_FALSE(timed_out) << "watchdog: the aborted fetch or the probe connection never finished"; + ASSERT_TRUE(completed); + EXPECT_EQ(resolver_calls, 1) << "the abort must have been issued from inside connect()"; + EXPECT_TRUE(fetch_error == asio::error::operation_aborted) << fetch_error.message(); + EXPECT_EQ(first_accepted_endpoint, probe_endpoint) + << "the exchange opened a connection after abort_pending() had already run: the listener " + "accepted it ahead of the probe"; +} + +#ifdef __linux__ + +// The window the test above pins, reached with the system resolver and close() called from a thread +// that does not run the io_context. Asio's resolver thread checks its cancel token once, before +// getaddrinfo(), so a close() arriving during that call can cancel neither the resolve nor a socket +// that is not open yet. The gate holds the lookup inside getaddrinfo() until the abort has run, so +// the flow resumes with usable addresses on a closed transport and must stop there; without the latch +// check in connect() it would block on the stalling server until the HTTP timeout. +TEST(AuthTransportCloseTest, CloseWhileTheResolverIsPastItsCancelCheckOpensNoConnection) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto stalling_base = stalling.base_url(); + + LoopbackServer resource_server(io_ctx); + const auto resource_base = resource_server.base_url(); + resource_server.set_handler([&](const http::request&) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer resource_metadata=")" + stalling_base + R"(/prm")"); + return challenge; + }); + asio::co_spawn(io_ctx, resource_server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = resource_base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {resource_base, stalling_base}; + policy.allow_plain_http_loopback = true; + config.policy = policy; + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + auto inner = + std::make_shared(io_ctx.get_executor(), resource_base + "/mcp"); + auto transport = std::make_shared(inner, manager); + + // Only the discovery fetch resolves the stalling server's port, so only that lookup is held. + resolve_gate().arm(stalling.port()); + + std::exception_ptr failure; + std::promise flow_done; + auto flow_finished = flow_done.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"tools/call"})"); + } catch (...) { + failure = std::current_exception(); + } + flow_done.set_value(); + }, + asio::detached); + + std::thread runner([&io_ctx]() { io_ctx.run(); }); + + const auto limit = std::chrono::seconds(10); + const bool lookup_held = resolve_gate().wait_until_entered(limit); + + // close() posts its abort to the io_context. Two further hops through the same queue cannot + // complete before it has, including when the abort had to wait its turn on the client's strand. + std::promise abort_done; + auto abort_finished = abort_done.get_future(); + bool abort_ran = false; + if (lookup_held) { + transport->close(); + asio::post(io_ctx, [&]() { asio::post(io_ctx, [&]() { abort_done.set_value(); }); }); + abort_ran = abort_finished.wait_for(limit) == std::future_status::ready; + } + + resolve_gate().release(); + const bool finished = flow_finished.wait_for(limit) == std::future_status::ready; + + // Everything below reads state the runner thread wrote, so it stops first. + io_ctx.stop(); + runner.join(); + + ASSERT_TRUE(lookup_held) << "the discovery fetch never reached getaddrinfo()"; + ASSERT_TRUE(abort_ran) << "the abort posted by close() never ran on the io_context"; + EXPECT_TRUE(finished) << "close() was lost: the flow resumed from the lookup after the " + "transport closed and is still running"; + EXPECT_EQ(stalling.accepted(), 0) + << "the flow connected to the metadata server after the transport had closed"; + if (finished) { + EXPECT_NE(failure, nullptr) << "a flow cut short by close() must report an error"; + } +} + +namespace { + +/// One request against a stalling server on an io_context run by two threads, with +/// abort_pending() called while one of them is inside the socket() call that opens the exchange's +/// connection. `request` issues it, and opens `sockets_before_the_held_one` connections first. +/// +/// The exchange has passed its last latch check, so only the posted close can stop it. Because the +/// exchange runs on the client's strand, the close runs once the exchange next suspends, with the +/// socket open; off the strand the second io thread would run it with nothing to close, and the +/// exchange would block on the stalling server until the HTTP timeout. +void expect_abort_during_socket_open_stops_the_exchange( + asio::io_context& io_ctx, const std::shared_ptr& client, + std::function()> request, int sockets_before_the_held_one = 0) { + socket_gate().arm(sockets_before_the_held_one); + + boost::system::error_code request_error; + std::promise request_done; + auto request_finished = request_done.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + co_await request(); + } catch (const boost::system::system_error& error) { + request_error = error.code(); + } catch (...) { + } + request_done.set_value(); + }, + asio::detached); + + std::thread first_runner([&io_ctx]() { io_ctx.run(); }); + std::thread second_runner([&io_ctx]() { io_ctx.run(); }); + + const auto limit = std::chrono::seconds(10); + const bool open_held = socket_gate().wait_until_entered(limit); + + // One io thread is parked in socket(), so the other runs everything posted from here, in + // order: a close that could run at all has run by the time the marker behind it does. + std::promise marker_done; + auto marker_finished = marker_done.get_future(); + bool marker_ran = false; + if (open_held) { + client->abort_pending(); + asio::post(io_ctx, [&]() { marker_done.set_value(); }); + marker_ran = marker_finished.wait_for(limit) == std::future_status::ready; + } + + socket_gate().release(); + const bool finished = request_finished.wait_for(limit) == std::future_status::ready; + + // Everything below reads state the runner threads wrote, so they stop first. + io_ctx.stop(); + first_runner.join(); + second_runner.join(); + + ASSERT_TRUE(open_held) << "the exchange never reached socket()"; + ASSERT_TRUE(marker_ran) << "the second io thread never ran the marker posted after the abort"; + ASSERT_TRUE(finished) << "abort_pending() was lost: its close ran while the exchange was " + "opening its socket on another thread, and the exchange is still " + "running"; + EXPECT_TRUE(request_error == asio::error::operation_aborted) << request_error.message(); +} + +} // namespace + +TEST(AuthTransportCloseTest, AbortWhileAnotherIoThreadOpensTheSocketStopsAGet) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto base = stalling.base_url(); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(base)); + + expect_abort_during_socket_open_stops_the_exchange( + io_ctx, client, [&]() -> mcp::Task { (void)co_await client->get_json(base + "/prm"); }); +} + +// A redirect is followed on a fresh socket, opened long after the exchange first reached the +// strand. The first hop's socket() passes the gate and the second hop's is held. +TEST(AuthTransportCloseTest, AbortWhileAnotherIoThreadOpensTheSocketStopsARedirectHop) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto stalling_base = stalling.base_url(); + + LoopbackServer redirecting(io_ctx); + const auto base = redirecting.base_url(); + redirecting.set_handler([&](const http::request&) { + http::response redirect{http::status::found, 11}; + redirect.set(http::field::location, stalling_base + "/prm"); + return redirect; + }); + asio::co_spawn(io_ctx, redirecting.serve(1), asio::detached); + + auto client = std::make_shared(io_ctx.get_executor()); + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {base, stalling_base}; + policy.allow_plain_http_loopback = true; + client->set_metadata_policy(policy); + + expect_abort_during_socket_open_stops_the_exchange( + io_ctx, client, [&]() -> mcp::Task { (void)co_await client->get_json(base + "/prm"); }, + 1); +} + +TEST(AuthTransportCloseTest, AbortWhileAnotherIoThreadOpensTheSocketStopsAJsonPost) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto base = stalling.base_url(); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(base)); + + const json registration = {{"client_name", "test"}}; + expect_abort_during_socket_open_stops_the_exchange(io_ctx, client, [&]() -> mcp::Task { + (void)co_await client->post_json(base + "/register", registration); + }); +} + +TEST(AuthTransportCloseTest, AbortWhileAnotherIoThreadOpensTheSocketStopsATokenRequest) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto base = stalling.base_url(); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(base)); + + mcp::auth::OAuthConfig config; + config.client_id = "test-client"; + config.token_endpoint = base + "/token"; + + expect_abort_during_socket_open_stops_the_exchange(io_ctx, client, [&]() -> mcp::Task { + (void)co_await client->refresh_token(config, "refresh-token"); + }); +} + +#endif // __linux__ + +// abort_pending() against exchanges at every stage of their life, on an io_context run by three +// threads and with nothing ordering the two: each abort is issued a varying few microseconds after +// its request starts. The exchange and the close abort_pending() posts touch the same socket, so +// they must never run at the same time. The assertions hold either way; the oracle for that +// property is ThreadSanitizer, which reports the two colliding unless both run on the client's +// strand. +TEST(AuthTransportCloseTest, AbortRacingExchangesOnThreeIoThreadsEndsEveryRequest) { + asio::io_context io_ctx; + auto work = asio::make_work_guard(io_ctx); + + // Accepts and immediately drops every connection, so an exchange the abort misses still ends. + asio::ip::tcp::acceptor acceptor(io_ctx, {asio::ip::make_address("127.0.0.1"), 0}); + const auto base = "http://127.0.0.1:" + std::to_string(acceptor.local_endpoint().port()); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + for (;;) { + boost::system::error_code accept_error; + auto socket = co_await acceptor.async_accept( + asio::redirect_error(asio::use_awaitable, accept_error)); + if (accept_error) { + co_return; + } + } + }, + asio::detached); + + std::vector runners; + for (int index = 0; index < 3; ++index) { + runners.emplace_back([&io_ctx]() { io_ctx.run(); }); + } + + const int requests = 2000; + int ended = 0; + for (int index = 0; index < requests; ++index) { + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(base)); + + std::promise request_done; + auto request_finished = request_done.get_future(); + asio::co_spawn( + io_ctx, + [&, client]() -> mcp::Task { + try { + (void)co_await client->get_json(base + "/prm"); + } catch (...) { + } + request_done.set_value(); + }, + asio::detached); + + const auto strike = + std::chrono::steady_clock::now() + std::chrono::microseconds((index * 13) % 400); + while (std::chrono::steady_clock::now() < strike) { + } + client->abort_pending(); + + if (request_finished.wait_for(std::chrono::seconds(10)) != std::future_status::ready) { + break; + } + ++ended; + } + + io_ctx.stop(); + for (auto& runner : runners) { + runner.join(); + } + + EXPECT_EQ(ended, requests) << "a request neither completed nor was aborted"; +} + +// A request runs on the client's strand, and its caller is resumed off it. A caller left on the +// strand would hold it for as long as it kept running, and a caller that blocks there -- here, on +// a gate the test holds -- would stall every other request the client has, along with the close +// abort_pending() posts, until it let go. +TEST(AuthHttpClientExecutorTest, ACallerThatBlocksAfterARequestDoesNotStallTheClient) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + server.set_handler( + [](const http::request&) { return json_response({{"ok", true}}); }); + asio::co_spawn(io_ctx, server.serve(2), asio::detached); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(base)); + + std::promise first_resumed; + auto first_has_resumed = first_resumed.get_future(); + std::promise release_first; + auto first_released = release_first.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await client->get_json(base + "/first"); + } catch (...) { + } + first_resumed.set_value(); + // Blocks the io thread it resumed on, as a caller doing synchronous work would. + (void)first_released.wait_for(std::chrono::seconds(30)); + }, + asio::detached); + + std::thread first_runner([&io_ctx]() { io_ctx.run(); }); + std::thread second_runner([&io_ctx]() { io_ctx.run(); }); + + const auto limit = std::chrono::seconds(10); + const bool resumed = first_has_resumed.wait_for(limit) == std::future_status::ready; + + std::promise second_done; + auto second_finished = second_done.get_future(); + bool finished = false; + if (resumed) { + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await client->get_json(base + "/second"); + } catch (...) { + } + second_done.set_value(); + }, + asio::detached); + finished = second_finished.wait_for(limit) == std::future_status::ready; + } + + release_first.set_value(); + io_ctx.stop(); + first_runner.join(); + second_runner.join(); + + ASSERT_TRUE(resumed) << "the first request never returned to its caller"; + EXPECT_TRUE(finished) << "a second request could not run while the first request's caller was " + "blocked: that caller was resumed on the client's strand"; +} + +#if BOOST_VERSION >= 107700 + +namespace { + +struct CancelledCallerOutcome { + bool resumed{false}; + bool has_result{false}; + boost::system::error_code error; + /// Whether the caller's later awaits still throw once it is cancelled, read after the request. + bool still_throws_if_cancelled{false}; + bool second_request_finished{false}; +}; + +/// Cancel a caller while its request is running, then let it block where it resumes and issue a +/// second request behind it. The cancellation is emitted from the resolver hook, which runs on the +/// client's strand in the middle of the exchange, so it lands at an exact place. +/// +/// The caller accepts every kind of cancellation; the exchange keeps the default and reacts to +/// terminal cancellation only. A terminal cancellation therefore ends the exchange, and a total one +/// marks the caller cancelled while the exchange runs to completion. +CancelledCallerOutcome cancel_a_caller_mid_request(asio::cancellation_type type) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + server.set_handler( + [](const http::request&) { return json_response({{"ok", true}}); }); + asio::co_spawn(io_ctx, server.serve(1), asio::detached); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(base)); + + asio::cancellation_signal cancel; + client->set_host_resolver([&cancel, type](const std::string&, const std::string&) { + cancel.emit(type); + return std::vector{"127.0.0.1"}; + }); + + CancelledCallerOutcome outcome; + std::promise first_resumed; + auto first_has_resumed = first_resumed.get_future(); + std::promise release_first; + auto first_released = release_first.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + co_await asio::this_coro::reset_cancellation_state(asio::enable_total_cancellation()); + try { + (void)co_await client->get_json(base + "/first"); + outcome.has_result = true; + } catch (const boost::system::system_error& error) { + outcome.error = error.code(); + } catch (...) { + } + outcome.still_throws_if_cancelled = co_await asio::this_coro::throw_if_cancelled(); + first_resumed.set_value(); + // Blocks the io thread it resumed on, as a caller doing synchronous work would. + (void)first_released.wait_for(std::chrono::seconds(30)); + }, + asio::bind_cancellation_slot(cancel.slot(), asio::detached)); + + std::thread first_runner([&io_ctx]() { io_ctx.run(); }); + std::thread second_runner([&io_ctx]() { io_ctx.run(); }); + + const auto limit = std::chrono::seconds(10); + outcome.resumed = first_has_resumed.wait_for(limit) == std::future_status::ready; + + std::promise second_done; + auto second_finished = second_done.get_future(); + if (outcome.resumed) { + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + // Outside the policy's origins: refused on the client's strand, with no + // lookup, so it ends at once unless something is holding that strand. + (void)co_await client->get_json("http://192.0.2.1/second"); + } catch (...) { + } + second_done.set_value(); + }, + asio::detached); + outcome.second_request_finished = second_finished.wait_for(limit) == std::future_status::ready; + } + + release_first.set_value(); + io_ctx.stop(); + first_runner.join(); + second_runner.join(); + return outcome; +} + +} // namespace + +// Cancelling the caller cancels its exchange, and the caller still leaves the client's strand +// before it runs again. co_await throws for a cancelled coroutine before initiating anything, so +// the post that takes the caller off the strand has to be made with that switched off. +TEST(AuthHttpClientExecutorTest, ACancelledCallerThatBlocksDoesNotStallTheClient) { + const auto outcome = cancel_a_caller_mid_request(asio::cancellation_type::terminal); + + ASSERT_TRUE(outcome.resumed) << "the cancelled request never returned to its caller"; + EXPECT_FALSE(outcome.has_result); + EXPECT_TRUE(outcome.error == asio::error::operation_aborted) << outcome.error.message(); + EXPECT_TRUE(outcome.still_throws_if_cancelled) + << "the request left the caller's coroutine ignoring its own cancellation"; + EXPECT_TRUE(outcome.second_request_finished) + << "a second request could not run while the cancelled caller was blocked: that caller " + "was resumed on the client's strand"; +} + +// A cancellation the exchange does not react to leaves it to finish, and what it fetched is still +// the caller's answer: leaving the strand must not turn a completed result into operation_aborted. +TEST(AuthHttpClientExecutorTest, AResultThatCompletedDespiteACancellationIsDelivered) { + const auto outcome = cancel_a_caller_mid_request(asio::cancellation_type::total); + + ASSERT_TRUE(outcome.resumed) << "the request never returned to its caller"; + EXPECT_TRUE(outcome.has_result) + << "the completed result was replaced by: " << outcome.error.message(); + EXPECT_TRUE(outcome.still_throws_if_cancelled) + << "the request left the caller's coroutine ignoring its own cancellation"; + EXPECT_TRUE(outcome.second_request_finished) + << "a second request could not run while the cancelled caller was blocked: that caller " + "was resumed on the client's strand"; +} + +#endif // BOOST_VERSION >= 107700 + +// The closed check and the flight read-or-create share one critical section in handle_challenge(). +// Split across two, a request that passed the check before close() ran could still create a fresh +// flight and run a full authorization flow after the transport had closed. One lock closes that +// window outright: a manager that is already closed refuses a new challenge before it does anything +// else, including the discovery fetch this asserts never happens. +TEST(AuthTransportCloseTest, CloseThenChallengeThrowsPromptlyWithoutAnyDiscoveryOrHttp) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + server.set_handler([&base](const http::request&) { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + + bool threw = false; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + manager->close(); + try { + (void)co_await manager->try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm")"); + } catch (const std::exception&) { + threw = true; + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(threw); + EXPECT_TRUE(server.targets().empty()); + EXPECT_TRUE(requested_scopes.empty()); +} + +// Authenticator::try_handle_challenge() is a coroutine on the virtual this overrides, so its contract +// is the lazy one: building the awaitable does nothing, and any error surfaces from the await. +// Impl::handle_challenge() contains no co_await or co_return, which makes it a plain function +// returning an awaitable, so a bare `throw` in its body would fire when try_handle_challenge() is +// *called* instead. A caller that builds the awaitable first and awaits it later -- or stores it, or +// hands it to a combinator -- would see the exception escape outside whatever try/catch was wrapped +// around the await. +TEST(AuthTransportCloseTest, ChallengeOnAClosedManagerThrowsFromTheAwaitNotFromTheCall) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + server.set_handler([&base](const http::request&) { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + }); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + + manager->close(); + // Well formed, so it is the closed check that refuses this and not the challenge parse. + const std::string header = R"(Bearer resource_metadata=")" + base + R"(/prm")"; + + std::optional> pending; + bool threw_from_the_call = false; + try { + pending.emplace(manager->try_handle_challenge(header)); + } catch (...) { + threw_from_the_call = true; + } + EXPECT_FALSE(threw_from_the_call); + ASSERT_TRUE(pending.has_value()); + + bool threw_from_the_await = false; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await std::move(*pending); + } catch (const std::exception&) { + threw_from_the_await = true; + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(threw_from_the_await); + EXPECT_TRUE(server.targets().empty()); + EXPECT_TRUE(requested_scopes.empty()); +} + +// Constructing a scope means dereferencing the client, so a null one faults in the constructor +// rather than later. Nonsensical usage either way; the point is that it is refused the way the +// sibling constructor refuses it, not that it segfaults. +TEST(AuthAuthenticatorConstructionTest, RefusesANullHttpClientInsteadOfFaulting) { + auto store = std::make_shared(); + mcp::auth::OAuthConfig config; + config.client_id = "test-client"; + + EXPECT_THROW(mcp::auth::OAuthAuthenticator(store, nullptr, config, "http://server1"), + std::invalid_argument); +} + +// The same property on OAuthAuthenticator. If close() were stateless there -- merely forwarding to +// OAuthHttpClient::abort_pending() and letting a request started afterward run -- then +// try_refresh_token() after close() would perform a real token-refresh exchange and overwrite the +// stored token. The sticky `aborted` flag on OAuthHttpClient closes this too: the refresh's own POST +// never opens a connection, run_refresh() folds that failure into its existing `false` return, and +// the stale token is left exactly as it was. +TEST(AuthTransportCloseTest, AuthenticatorCloseThenRefreshPerformsNoNetworkIO) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + int token_requests = 0; + server.set_handler([&](const http::request&) { + ++token_requests; + return json_response(token_document()); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + auto store = std::make_shared(); + mcp::auth::TokenResponse stored; + stored.access_token = "stale-access-token"; + stored.token_type = "Bearer"; + stored.refresh_token = "stale-refresh-token"; + store->store(base + "/mcp", stored); + + mcp::auth::OAuthConfig config; + config.client_id = "test-client"; + config.token_endpoint = base + "/token"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + + auto http_client = std::make_shared(io_ctx.get_executor()); + http_client->set_metadata_policy(loopback_policy(server.origin())); + auto authenticator = + std::make_shared(store, http_client, config, base + "/mcp"); + + bool refreshed = true; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + authenticator->close(); + refreshed = co_await authenticator->try_refresh_token(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_FALSE(refreshed); + EXPECT_EQ(token_requests, 0); + EXPECT_TRUE(server.targets().empty()); + // The stale token was left alone, not overwritten by a refresh that should never have run. + EXPECT_EQ(authenticator->get_access_token(), "stale-access-token"); +} + +// Closing one authenticator must not disable another that shares the same HTTP client, and a churn of +// short-lived authenticators must leave no per-scope record behind on the client. +// +// LIMITATION: this test cannot fail on its own. The accessor answers zero structurally, so a newly +// introduced per-scope container would not be detected here. +// TheScopeAbortLatchIsReleasedByEveryChurnedAuthenticator below is the live guard for scope lifetime +// -- edit that one. +TEST(AuthHttpClientScopeRetentionTest, RetainsNoPerScopeStateAsAuthenticatorsComeAndGo) { + asio::io_context io_ctx; + auto store = std::make_shared(); + auto http_client = std::make_shared(io_ctx.get_executor()); + + mcp::auth::OAuthConfig config; + config.client_id = "churn-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + + ASSERT_EQ(mcp::auth::internal::retained_scope_record_count(*http_client), 0U) + << "a client that has issued no scope at all is already holding records"; + + // Each authenticator opens a scope, closes it, and is destroyed before the next one is built, + // so at no point are two of them alive together. Nothing here is still reachable afterward. + const auto churn_through = [&](int count, int first_index) { + for (int index = 0; index < count; ++index) { + mcp::auth::OAuthAuthenticator authenticator( + store, http_client, config, "http://server" + std::to_string(first_index + index)); + authenticator.close(); + } + }; + + constexpr int small_churn = 8; + constexpr int large_churn = 256; + churn_through(small_churn, 0); + const auto after_small = mcp::auth::internal::retained_scope_record_count(*http_client); + churn_through(large_churn, small_churn); + const auto after_large = mcp::auth::internal::retained_scope_record_count(*http_client); + + // The shape of the bug is proportionality: what the client keeps must not be a function of how + // many authenticators have come and gone. + EXPECT_EQ(after_small, after_large) + << "retained records grew from " << after_small << " to " << after_large << " over " + << large_churn << " more closed authenticators"; + EXPECT_EQ(after_large, 0U) << "the client kept " << after_large + << " abort records for authenticators that are gone"; +} + +// The live guard for scope-latch lifetime, and the one to edit if you change how scopes are held. +// +// Counts the abort latches alive in the process, drives a churn of authenticators that each run a +// real token refresh through their scope, and requires the count to return to its starting value. The +// peak assertion keeps the test from passing vacuously against instrumentation that always reports +// zero; the peak is the number of authenticators the test is holding. +TEST(AuthHttpClientScopeRetentionTest, TheScopeAbortLatchIsReleasedByEveryChurnedAuthenticator) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + constexpr int churn = 16; + + auto store = std::make_shared(); + auto http_client = std::make_shared(io_ctx.get_executor()); + http_client->set_metadata_policy(loopback_policy(server.origin())); + + const auto baseline = mcp::auth::internal::live_scope_latch_count(); + std::size_t peak = baseline; + int in_flight_samples = 0; + int refreshed = 0; + + // Single-threaded io_context: the handler runs on the same thread as the churn coroutine, so + // these are plain variables rather than atomics. + server.set_handler([&](const http::request&) { + ++in_flight_samples; + peak = std::max(peak, mcp::auth::internal::live_scope_latch_count()); + return json_response(token_document()); + }); + asio::co_spawn(io_ctx, server.serve(churn), asio::detached); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + // Held together on purpose: the peak has to exceed the baseline by more than one, so a + // counter that merely toggled would not satisfy it. + std::vector> live; + for (int index = 0; index < churn; ++index) { + const auto server_url = base + "/server" + std::to_string(index); + mcp::auth::TokenResponse stored; + stored.access_token = "stale-access-token"; + stored.token_type = "Bearer"; + stored.refresh_token = "stale-refresh-token"; + store->store(server_url, stored); + + mcp::auth::OAuthConfig config; + config.client_id = "churn-client"; + config.token_endpoint = base + "/token"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + + auto authenticator = std::make_shared( + store, http_client, config, server_url); + // A real exchange through the scope: the token server's handler runs inside this + // co_await, and that is where the peak is sampled. + if (co_await authenticator->try_refresh_token()) { + ++refreshed; + } + authenticator->close(); + live.push_back(std::move(authenticator)); + } + // Every scope goes away here; nothing the test holds references a latch afterward. + live.clear(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_EQ(refreshed, churn) << "the churn did not actually perform its exchanges"; + ASSERT_EQ(in_flight_samples, churn) + << "the peak was never sampled with an exchange in flight, so it reports the count at some " + "other moment and the comment above is describing something the test does not measure"; + EXPECT_GT(peak, baseline + 1) + << "latches were never counted, so returning to the baseline proves nothing"; + EXPECT_EQ(mcp::auth::internal::live_scope_latch_count(), baseline) + << "a scope's abort latch outlived the authenticator that owned it"; +} + +// OAuthAuthenticator takes its client by shared_ptr, which is an invitation to share one across +// several servers, so close() must end only that authenticator's own scope. Latching the whole +// client instead would disable every other holder, and silently: run_refresh() reports a failed +// refresh as a plain `false`, so the survivor raises nothing and an application sees tokens +// quietly stop renewing against a server it never closed. +TEST(AuthTransportCloseTest, ClosingOneAuthenticatorLeavesAnotherSharingTheSameClientWorking) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + std::vector token_targets; + server.set_handler([&](const http::request& request) { + token_targets.emplace_back(request.target()); + return json_response(token_document()); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + auto store = std::make_shared(); + const auto closed_server = base + "/closed-server"; + const auto surviving_server = base + "/surviving-server"; + for (const auto& server_url : {closed_server, surviving_server}) { + mcp::auth::TokenResponse stored; + stored.access_token = "stale-access-token"; + stored.token_type = "Bearer"; + stored.refresh_token = "stale-refresh-token"; + store->store(server_url, stored); + } + + // One client, shared. This is the shape the constructor's signature invites. + auto http_client = std::make_shared(io_ctx.get_executor()); + http_client->set_metadata_policy(loopback_policy(server.origin())); + + mcp::auth::OAuthConfig closed_config; + closed_config.client_id = "closed-client"; + closed_config.token_endpoint = base + "/closed-token"; + closed_config.redirect_uri = "http://127.0.0.1:9999/callback"; + + mcp::auth::OAuthConfig surviving_config; + surviving_config.client_id = "surviving-client"; + surviving_config.token_endpoint = base + "/surviving-token"; + surviving_config.redirect_uri = "http://127.0.0.1:9999/callback"; + + auto closing = std::make_shared(store, http_client, closed_config, + closed_server); + auto surviving = std::make_shared( + store, http_client, surviving_config, surviving_server); + + bool closed_refreshed = true; + bool surviving_refreshed = false; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + closing->close(); + // The closed one is still closed: its own refresh must not reach the network. + closed_refreshed = co_await closing->try_refresh_token(); + // The one nobody closed must be entirely unaffected. + surviving_refreshed = co_await surviving->try_refresh_token(); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_FALSE(closed_refreshed); + EXPECT_TRUE(surviving_refreshed) + << "closing one authenticator disabled another that only shares the client"; + // Exactly one token request, from the surviving authenticator, at its own endpoint. + const std::vector expected_targets{"/surviving-token"}; + EXPECT_EQ(token_targets, expected_targets); + EXPECT_EQ(closing->get_access_token(), "stale-access-token"); + EXPECT_EQ(surviving->get_access_token(), "granted-access-token"); +} + +// close() must be safe to call more than once (OAuthClientTransport::close() itself is idempotent +// and calls it only once per transport, but the manager and authenticator are reachable directly), +// and a write issued after close() must fail fast rather than hang or reach the network. +TEST(AuthTransportCloseTest, CloseTwiceIsIdempotentAndARequestAfterCloseFailsFast) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + server.set_handler([&base](const http::request&) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer resource_metadata=")" + base + R"(/prm")"); + return challenge; + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + auto inner = std::make_shared(io_ctx.get_executor(), base + "/mcp"); + auto transport = std::make_shared(inner, manager); + + bool completed = false; + std::exception_ptr failure; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + transport->close(); + transport->close(); // Must not throw, hang, or double-release anything. + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"tools/call"})"); + } catch (...) { + failure = std::current_exception(); + } + completed = true; + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(completed); + EXPECT_NE(failure, nullptr); + EXPECT_TRUE(server.targets().empty()); +} + +TEST(AuthTransportCloseTest, CloseWithQueuedPendingRequestsFailsThemPromptlyWithoutHanging) { + asio::io_context io_ctx; + StallingServer stalling(io_ctx); + stalling.accept_and_stall(); + const auto stalling_base = stalling.base_url(); + + LoopbackServer resource_server(io_ctx); + const auto resource_base = resource_server.base_url(); + resource_server.set_handler([&](const http::request&) { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer resource_metadata=")" + stalling_base + R"(/prm")"); + return challenge; + }); + asio::co_spawn(io_ctx, resource_server.serve(10), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = resource_base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins = {resource_base, stalling_base}; + policy.allow_plain_http_loopback = true; + config.policy = policy; + + auto store = std::make_shared(); + std::vector requested_scopes; + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + auto inner = + std::make_shared(io_ctx.get_executor(), resource_base + "/mcp"); + auto transport = std::make_shared(inner, manager); + + constexpr int request_count = 3; + std::vector wires; + for (int index = 0; index < request_count; ++index) { + wires.push_back(json{{"jsonrpc", "2.0"}, {"id", index}, {"method", "tools/call"}}.dump()); + } + + int completed = 0; + int failed = 0; + bool timed_out = false; + + asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(std::chrono::seconds(10)); + watchdog.async_wait([&](boost::system::error_code error) { + if (!error && completed < request_count) { + timed_out = true; + io_ctx.stop(); + } + }); + + for (const auto& wire : wires) { + asio::co_spawn( + io_ctx, + [&transport, &failed, &completed, &watchdog, wire]() -> mcp::Task { + try { + co_await transport->write_message(wire); + } catch (...) { + ++failed; + } + ++completed; + if (completed == request_count) { + watchdog.cancel(); + } + }, + asio::detached); + } + + asio::steady_timer closer(io_ctx); + closer.expires_after(std::chrono::milliseconds(200)); + closer.async_wait([&](boost::system::error_code) { + transport->close(); + resource_server.close(); + }); + + io_ctx.run(); + + ASSERT_FALSE(timed_out) << "watchdog: close() left a queued request parked"; + EXPECT_EQ(completed, request_count); + EXPECT_EQ(failed, request_count); +} + +TEST(AuthIssuerBindingTest, RejectsAuthServerMetadataWhoseIssuerIsNotItsDiscoveryLocation) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/mcp") { + http::response challenge{http::status::unauthorized, 11}; + challenge.set(http::field::www_authenticate, + R"(Bearer resource_metadata=")" + base + R"(/prm")"); + return challenge; + } + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + // The document claims to be a different authorization server than the one it was + // fetched from -- RFC 8414 3.3 makes that a hard refusal, not a normalization problem. + auto metadata = auth_server_metadata(base, true); + metadata["issuer"] = "https://other.example.com"; + return json_response(metadata); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(5), asio::detached); + + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + auto store = std::make_shared(); + std::vector requested_scopes; + bool threw = false; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + auto manager = std::make_shared( + io_ctx.get_executor(), store, config, recording_callback(&requested_scopes)); + try { + (void)co_await manager->try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm")"); + } catch (const std::exception&) { + threw = true; + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(threw); + // Refused before any consent prompt or token exchange could bind a credential to the wrong AS. + EXPECT_TRUE(requested_scopes.empty()); + EXPECT_FALSE(store->load(config.server_url).has_value()); +} + +TEST(AuthClientIdentityBindingTest, RefusesInjectedCredentialsBoundToADifferentIssuer) { + mcp::auth::ClientIdentityConfig config; + mcp::auth::OAuthClientInformation injected; + injected.client_id = "client-for-as-one"; + injected.client_secret = "secret-for-as-one"; + injected.issuer = "https://as-one.example.com"; + config.pre_registered = injected; + + mcp::auth::ClientIdentityServerFacts other; + other.issuer = "https://as-two.example.com"; + // A registration endpoint is on offer, and it must still not be taken: injected credentials are + // terminal, so a misbound secret yields no identity rather than a silent dynamic registration. + other.registration_endpoint = "https://as-two.example.com/register"; + + EXPECT_EQ(mcp::auth::select_client_identity(config, other, std::nullopt), + mcp::auth::ClientIdentityDecision::unavailable); + + mcp::auth::ClientIdentityServerFacts matching; + matching.issuer = "https://as-one.example.com"; + EXPECT_EQ(mcp::auth::select_client_identity(config, matching, std::nullopt), + mcp::auth::ClientIdentityDecision::use_pre_registered); + + // Credentials that never named an issuer are unbound, not universally bound. A secret that + // belongs to nobody in particular must not be handed to whichever authorization server the + // protected-resource document happened to name. + config.pre_registered->issuer.clear(); + EXPECT_EQ(mcp::auth::select_client_identity(config, other, std::nullopt), + mcp::auth::ClientIdentityDecision::unavailable); + EXPECT_EQ(mcp::auth::select_client_identity(config, matching, std::nullopt), + mcp::auth::ClientIdentityDecision::unavailable); +} + +// The refusal is aimed at the secret, not at injected credentials in general. A public client's +// `client_id` is not confidential, so an unbound one still authorizes normally. +TEST(AuthClientIdentityBindingTest, UnboundPublicClientCredentialsAreStillUsed) { + mcp::auth::ClientIdentityConfig config; + mcp::auth::OAuthClientInformation injected; + injected.client_id = "public-client"; + config.pre_registered = injected; + + mcp::auth::ClientIdentityServerFacts anywhere; + anywhere.issuer = "https://as-two.example.com"; + anywhere.registration_endpoint = "https://as-two.example.com/register"; + + EXPECT_EQ(mcp::auth::select_client_identity(config, anywhere, std::nullopt), + mcp::auth::ClientIdentityDecision::use_pre_registered); + + // An empty-string secret is no secret at all and must not trip the refusal. + config.pre_registered->client_secret = ""; + EXPECT_EQ(mcp::auth::select_client_identity(config, anywhere, std::nullopt), + mcp::auth::ClientIdentityDecision::use_pre_registered); +} + +namespace { + +/// What an end-to-end run of the shorthand `client_id` / `client_secret` / `client_issuer` config +/// did, seen from the authorization server's side of the wire. +struct InjectedSecretOutcome { + bool authorized{false}; + std::string failure; + std::vector targets; + bool secret_seen_on_the_wire{false}; +}; + +/// Drive one challenge with shorthand credentials whose bound issuer is `choose_issuer(base)`, +/// against a loopback authorization server that advertises `client_secret_post` so a presented +/// secret lands in the token request body verbatim and can be asserted on directly. +/// +/// The issuer is chosen from the server's base URL rather than passed in, because the fixture binds +/// an ephemeral port that the caller cannot know before the server exists. +InjectedSecretOutcome try_injected_secret( + const std::function& choose_issuer) { + static constexpr std::string_view secret = "application-held-secret"; + + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + auto metadata = auth_server_metadata(base, true); + metadata["token_endpoint_auth_methods_supported"] = json::array({"client_secret_post"}); + return json_response(metadata); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(3), asio::detached); + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "application-held-client"; + config.client_secret = std::string(secret); + config.client_issuer = choose_issuer(base); + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + // asio::detached swallows exceptions, which would turn a failed assertion into a hang. + std::promise result; + auto observed = result.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + InjectedSecretOutcome outcome; + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + try { + outcome.authorized = co_await manager.try_handle_challenge( + R"(Bearer resource_metadata=")" + base + R"(/prm")"); + } catch (const std::exception& error) { + outcome.failure = error.what(); + } + outcome.targets = server.targets(); + // The secret would travel in the token request body under client_secret_post, and in + // the base64 Authorization header under client_secret_basic; check both spellings so + // this cannot pass merely because the auth method changed. + const std::string basic_credentials = "application-held-client:" + std::string(secret); + const auto basic_header = mcp::auth::detail::base64_encode( + reinterpret_cast(basic_credentials.data()), + basic_credentials.size()); + const auto contains_secret = [&](const std::vector& recorded) { + return std::any_of(recorded.begin(), recorded.end(), [&](const std::string& value) { + return value.find(secret) != std::string::npos || + value.find(basic_header) != std::string::npos; + }); + }; + outcome.secret_seen_on_the_wire = + contains_secret(server.bodies()) || contains_secret(server.authorizations()); + result.set_value(std::move(outcome)); + server.close(); + }, + asio::detached); + + io_ctx.run(); + return observed.get(); +} + +bool saw_target(const std::vector& targets, std::string_view target) { + return std::find(targets.begin(), targets.end(), target) != targets.end(); +} + +} // namespace + +// The invariant: a client_secret never reaches an authorization server it is not bound to. The +// protected-resource document names this authorization server, so binding the credentials to a +// different one must stop the flow before the token request, not merely fail it afterwards. +TEST(AuthClientIdentityBindingTest, DoesNotSendTheClientSecretToAnAuthorizationServerItIsNotBoundTo) { + const auto outcome = + try_injected_secret([](const std::string&) { return "https://as-elsewhere.example.com"; }); + + EXPECT_FALSE(outcome.authorized); + EXPECT_FALSE(outcome.failure.empty()); + EXPECT_FALSE(saw_target(outcome.targets, "/token")); + EXPECT_FALSE(outcome.secret_seen_on_the_wire); +} + +// Credentials that name no issuer must not fall through to `use_pre_registered` for every +// authorization server, which would leave the guard inert on the shorthand path the SDK itself +// builds. Refusing at the point of use keeps construction working for every caller while still +// guaranteeing the secret is never transmitted. +TEST(AuthClientIdentityBindingTest, RefusesAnInjectedClientSecretThatNamesNoIssuer) { + const auto outcome = try_injected_secret([](const std::string&) { return std::string{}; }); + + EXPECT_FALSE(outcome.authorized); + EXPECT_NE(outcome.failure.find("name no issuer"), std::string::npos) << outcome.failure; + EXPECT_FALSE(saw_target(outcome.targets, "/token")); + EXPECT_FALSE(outcome.secret_seen_on_the_wire); +} + +// The positive control, without which the two tests above would pass on a build that simply never +// authorizes. A correctly bound secret still reaches the token endpoint it belongs to, and this is +// also what proves the wire assertions above can observe a secret when one is really sent. +TEST(AuthClientIdentityBindingTest, AuthorizesNormallyWhenTheInjectedSecretNamesItsOwnIssuer) { + const auto outcome = try_injected_secret([](const std::string& base) { return base; }); + + EXPECT_TRUE(outcome.authorized) << outcome.failure; + EXPECT_TRUE(outcome.failure.empty()) << outcome.failure; + EXPECT_TRUE(saw_target(outcome.targets, "/token")); + EXPECT_TRUE(outcome.secret_seen_on_the_wire); +} + +// `client_issuer` does not clear the refusal above for hand-built credentials, and the refusal has to +// say so: the manager's constructor copies `client_issuer` into the injected credentials only when +// `client_identity.pre_registered` was not already populated. Pins that a hand-built `pre_registered` +// stays refused with `client_issuer` set, and that the message names the field that would work. +TEST(AuthClientIdentityBindingTest, RefusalNamesTheFieldThatAppliesToHandBuiltCredentials) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + auto metadata = auth_server_metadata(base, true); + metadata["token_endpoint_auth_methods_supported"] = json::array({"client_secret_post"}); + return json_response(metadata); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(3), asio::detached); + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + // Hand-built rather than the client_id/client_secret shorthand: this is the path on which + // client_issuer is ignored. + mcp::auth::OAuthClientInformation injected; + injected.client_id = "application-held-client"; + injected.client_secret = "application-held-secret"; + injected.issuer.clear(); + config.client_identity.pre_registered = injected; + + // Set to the RIGHT issuer, and deliberately so: if `client_issuer` were the remedy here, the + // flow below would authorize. It does not, because nothing reads it here. + config.client_issuer = base; + + std::promise> result; + auto observed = result.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + bool authorized = false; + std::string failure; + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + try { + authorized = co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + + base + R"(/prm")"); + } catch (const std::exception& error) { + failure = error.what(); + } + result.set_value({authorized, failure}); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + const auto [authorized, failure] = observed.get(); + EXPECT_FALSE(authorized) << "client_issuer is not read on this path, so the refusal must stand"; + ASSERT_FALSE(failure.empty()); + EXPECT_NE(failure.find("name no issuer"), std::string::npos) << failure; + // The remedy that actually applies here. + EXPECT_NE(failure.find("client_identity.pre_registered.issuer"), std::string::npos) << failure; + // And the message must say plainly that the field the caller already set does nothing here, + // rather than listing it as an equal alternative. + EXPECT_NE(failure.find("ignored once"), std::string::npos) << failure; + // No token request may have been attempted. + EXPECT_FALSE(saw_target(server.targets(), "/token")); +} + +namespace { + +/// Drive a challenge whose protected-resource metadata names `prm_resource` against a manager +/// configured with `server_url`, and report whether authorization completed and, when it did, the +/// `resource` value carried on the request that reached the authorization server. +/// +/// `server_url` need not be reachable: the challenge always names the metadata location explicitly, +/// so discovery never derives a fetch target from it. Only the `resource` comparison in +/// RFC 9728 §3.3 validation reads it. +struct PrmResourceOutcome { + bool authorized{false}; + bool threw{false}; + std::string resolved_resource; +}; + +PrmResourceOutcome try_prm_resource(const std::string& server_url, const std::string& prm_resource) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm.json") { + return json_response( + {{"resource", prm_resource}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(3), asio::detached); + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = server_url; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + PrmResourceOutcome outcome; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + try { + outcome.authorized = co_await manager.try_handle_challenge( + R"(Bearer realm="mcp", resource_metadata=")" + base + R"(/prm.json")"); + if (const auto record = manager.last_authorization_request(); + record && record->resource) { + outcome.resolved_resource = *record->resource; + } + } catch (const std::exception&) { + outcome.threw = true; + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + return outcome; +} + +} // namespace + +// `server_url` and the PRM `resource` value are only ever compared, never fetched, so a synthetic +// origin exercises the comparison logic without any network dependency. + +// RFC 9728 §3.3: byte-exact match is always accepted, and the PRM's own value is what travels on +// the authorization request (not a value substituted from config). +TEST(AuthProtectedResourceValidationTest, AcceptsAByteExactMatch) { + const auto outcome = try_prm_resource("https://example.test/mcp", "https://example.test/mcp"); + EXPECT_FALSE(outcome.threw); + EXPECT_TRUE(outcome.authorized); + EXPECT_EQ(outcome.resolved_resource, "https://example.test/mcp"); +} + +// The root-PRM layout (conformance `auth/metadata-var2`): the PRM's `resource` legitimately +// identifies the server at coarser granularity than the endpoint URL. An origin-only value must be +// accepted for any path on that origin. +TEST(AuthProtectedResourceValidationTest, AcceptsAnOriginOnlyResourceForAPathedServerUrl) { + const auto outcome = try_prm_resource("https://example.test/mcp", "https://example.test"); + EXPECT_FALSE(outcome.threw); + EXPECT_TRUE(outcome.authorized); + // The PRM's own (coarser) value is what is used, not the finer server URL. + EXPECT_EQ(outcome.resolved_resource, "https://example.test"); +} + +TEST(AuthProtectedResourceValidationTest, AcceptsAPathPrefixAlignedOnASegmentBoundary) { + const auto outcome = try_prm_resource("https://example.test/a/b", "https://example.test/a"); + EXPECT_FALSE(outcome.threw); + EXPECT_TRUE(outcome.authorized); + EXPECT_EQ(outcome.resolved_resource, "https://example.test/a"); +} + +// A single trailing `/` on the PRM's `resource` path is ignored, so a resource value that trails +// its path with `/` is accepted against a server URL that does not. +TEST(AuthProtectedResourceValidationTest, AcceptsATrailingSlashOnTheResourceAgainstAnEqualPath) { + const auto outcome = try_prm_resource("https://example.test/a/b", "https://example.test/a/b/"); + EXPECT_FALSE(outcome.threw); + EXPECT_TRUE(outcome.authorized); + // The PRM's own value (with its trailing `/`) is what travels on the request. + EXPECT_EQ(outcome.resolved_resource, "https://example.test/a/b/"); +} + +// The mirror of the above: a server URL that trails its path with `/` was already accepted against +// a resource value that does not, and must remain so. +TEST(AuthProtectedResourceValidationTest, AcceptsATrailingSlashOnTheServerAgainstAnEqualPath) { + const auto outcome = try_prm_resource("https://example.test/a/b/", "https://example.test/a/b"); + EXPECT_FALSE(outcome.threw); + EXPECT_TRUE(outcome.authorized); + EXPECT_EQ(outcome.resolved_resource, "https://example.test/a/b"); +} + +// Stripping the resource's trailing `/` happens before the segment-boundary-prefix comparison too: +// "/a/" becomes "/a", which is a boundary-aligned prefix of "/a/b". +TEST(AuthProtectedResourceValidationTest, AcceptsATrailingSlashResourceAsABoundaryPrefix) { + const auto outcome = try_prm_resource("https://example.test/a/b", "https://example.test/a/"); + EXPECT_FALSE(outcome.threw); + EXPECT_TRUE(outcome.authorized); + EXPECT_EQ(outcome.resolved_resource, "https://example.test/a/"); +} + +// The lone "/" (origin-root) resource value is a distinct case from a trailing slash on a +// non-empty path, and its any-path-on-this-origin acceptance is unaffected by the trailing-slash +// stripping above. +TEST(AuthProtectedResourceValidationTest, AcceptsALoneSlashResourceForAPathedServerUrl) { + const auto outcome = try_prm_resource("https://example.test/mcp", "https://example.test/"); + EXPECT_FALSE(outcome.threw); + EXPECT_TRUE(outcome.authorized); + EXPECT_EQ(outcome.resolved_resource, "https://example.test/"); +} + +// "/ap" textually prefixes "/api", but not on a `/` segment boundary, so it must not be accepted as +// identifying it. +TEST(AuthProtectedResourceValidationTest, RejectsANonBoundaryPathPrefix) { + const auto outcome = try_prm_resource("https://example.test/api", "https://example.test/ap"); + EXPECT_TRUE(outcome.threw); + EXPECT_FALSE(outcome.authorized); +} + +TEST(AuthProtectedResourceValidationTest, RejectsADifferentAuthority) { + const auto outcome = + try_prm_resource("https://example.test/mcp", "https://different.example.test/mcp"); + EXPECT_TRUE(outcome.threw); + EXPECT_FALSE(outcome.authorized); +} + +TEST(AuthProtectedResourceValidationTest, RejectsADifferentScheme) { + const auto outcome = try_prm_resource("https://example.test/mcp", "http://example.test/mcp"); + EXPECT_TRUE(outcome.threw); + EXPECT_FALSE(outcome.authorized); +} + +// Authority comparison is byte-exact with no normalization: an explicit default port is a +// different authority than an implicit one, even though they denote the same origin. +TEST(AuthProtectedResourceValidationTest, RejectsAnExplicitDefaultPortAgainstAnImplicitOne) { + const auto outcome = try_prm_resource("https://example.test:443/mcp", "https://example.test/mcp"); + EXPECT_TRUE(outcome.threw); + EXPECT_FALSE(outcome.authorized); +} + +// A query or fragment on the PRM's `resource` value is never allowed, even when the origin and path +// would otherwise match. +TEST(AuthProtectedResourceValidationTest, RejectsAResourceValueCarryingAQuery) { + asio::io_context probe; + LoopbackServer server(probe); + const auto base = server.base_url(); + server.close(); + + const auto outcome = try_prm_resource(base + "/mcp", base + "/mcp?tenant=1"); + EXPECT_TRUE(outcome.threw); + EXPECT_FALSE(outcome.authorized); +} + +// RFC 9728 §2 makes `resource` a required PRM member. An empty value -- indistinguishable, once +// parsed, from a PRM document that omits the member entirely -- must be rejected rather than +// silently falling back to the configured server URL. +TEST(AuthProtectedResourceValidationTest, RejectsAnEmptyPrmResource) { + const auto outcome = try_prm_resource("https://example.test/mcp", ""); + EXPECT_TRUE(outcome.threw); + EXPECT_FALSE(outcome.authorized); +} + +namespace { + +/// Credential store that only counts, so a test can assert nothing was ever persisted. +class CountingCredentialStore final : public mcp::auth::ClientCredentialStore { + public: + void store(const std::string& issuer, mcp::auth::OAuthClientInformation information) override { + (void)issuer; + (void)information; + ++stores; + } + + [[nodiscard]] std::optional load( + const std::string& issuer) const override { + (void)issuer; + return std::nullopt; + } + + void remove(const std::string& issuer) override { (void)issuer; } + + int stores{0}; +}; + +} // namespace + +// Rejecting an unidentified `resource` is not enough on its own: it has to happen before the SDK +// acts on anything else the same untrusted document names. A PRM that points at an authorization +// server we would otherwise fetch metadata from, register a client with, and persist credentials +// for must cause none of those, so the only request this fixture ever sees is the PRM fetch itself. +TEST(AuthProtectedResourceValidationTest, + RejectsABadResourceBeforeDiscoveringTheAuthorizationServerOrRegisteringAClient) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm.json") { + // `resource` names a different origin entirely, so it does not identify our server. + return json_response({{"resource", "https://attacker.test/mcp"}, + {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + auto metadata = auth_server_metadata(base, true); + metadata["registration_endpoint"] = base + "/register"; + return json_response(metadata); + } + if (target == "/register") { + return json_response( + {{"client_id", "attacker-minted-client"}, {"client_secret", "attacker-minted-secret"}}); + } + if (target == "/token") { + return json_response(token_document()); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(4), asio::detached); + + auto credentials = std::make_shared(); + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + // No client_id, so identity resolution would reach dynamic registration. + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.credential_store = credentials; + config.policy = loopback_policy(server.origin()); + + // asio::detached swallows exceptions, which would turn a failed assertion into a hang; the + // promise carries the outcome back to the test body instead. + std::promise failure; + auto observed = failure.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + std::string message; + try { + (void)co_await manager.try_handle_challenge( + R"(Bearer realm="mcp", resource_metadata=")" + base + R"(/prm.json")"); + message = ""; + } catch (const std::exception& error) { + message = error.what(); + } + failure.set_value(message); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + const auto message = observed.get(); + EXPECT_NE(message.find("does not identify server"), std::string::npos) << message; + // The PRM fetch is the only request that may have happened: no authorization-server metadata + // discovery, and no dynamic-registration POST. + const std::vector expected_targets{"/prm.json"}; + EXPECT_EQ(server.targets(), expected_targets); + EXPECT_EQ(credentials->stores, 0); +} + +// Rejecting the document is only half of it: it must not be cached either. A document that merely +// parses must not reach the resource cache ahead of the identity check, or a refused document sits +// planted under the challenge's own metadata URL for the full TTL and the next attempt is served +// it without touching the network -- which is what makes an attacker-supplied document worth +// planting. Discovery takes the caller's acceptance test and commits nothing the caller refuses, +// so the request log is the evidence: two PRM fetches, not one. +TEST(AuthProtectedResourceValidationTest, ARejectedProtectedResourceDocumentIsNotCached) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm.json") { + // `resource` names a different origin, so it does not identify our server. + return json_response({{"resource", "https://attacker.test/mcp"}, + {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(4), asio::detached); + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + // asio::detached swallows exceptions, which would turn a failed assertion into a hang; the + // promise carries both outcomes back to the test body instead. + std::promise> failures; + auto observed = failures.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + // One manager for both attempts: the discovery cache it owns is the thing under test, + // and a second manager would have an empty one for trivial reasons. + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + const std::string header = + R"(Bearer realm="mcp", resource_metadata=")" + base + R"(/prm.json")"; + + std::pair messages; + try { + (void)co_await manager.try_handle_challenge(header); + messages.first = ""; + } catch (const std::exception& error) { + messages.first = error.what(); + } + try { + (void)co_await manager.try_handle_challenge(header); + messages.second = ""; + } catch (const std::exception& error) { + messages.second = error.what(); + } + failures.set_value(messages); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + const auto messages = observed.get(); + EXPECT_NE(messages.first.find("does not identify server"), std::string::npos) << messages.first; + // The second attempt must fail for the same reason, and must have re-fetched to find out. + EXPECT_NE(messages.second.find("does not identify server"), std::string::npos) << messages.second; + const std::vector expected_targets{"/prm.json", "/prm.json"}; + EXPECT_EQ(server.targets(), expected_targets) + << "the rejected document was served from the cache instead of being re-fetched"; +} + +namespace { + +/// A payload shaped like a forged log entry: the CR/LF closes the SDK's own line, and what follows +/// reads as a fresh, authoritative-looking one. +const std::string& forged_log_line() { + static const std::string value = + "\r\n2026-09-20T00:00:00Z INFO authorization granted to everyone\r\n"; + return value; +} + +bool carries_control_characters(const std::string& text) { + return std::any_of(text.begin(), text.end(), [](char character) { + const auto value = static_cast(character); + return value < 0x20 || value == 0x7f; + }); +} + +/// Whether `text` is well-formed UTF-8. Deliberately written out rather than delegated to the JSON +/// library, so the assertion does not depend on the same code the SDK might be using. +bool is_well_formed_utf8(const std::string& text) { + std::size_t index = 0; + while (index < text.size()) { + const auto lead = static_cast(text[index]); + std::size_t length = 0; + std::uint32_t codepoint = 0; + if (lead < 0x80) { + ++index; + continue; + } + if ((lead & 0xE0) == 0xC0) { + length = 2; + codepoint = lead & 0x1FU; + } else if ((lead & 0xF0) == 0xE0) { + length = 3; + codepoint = lead & 0x0FU; + } else if ((lead & 0xF8) == 0xF0) { + length = 4; + codepoint = lead & 0x07U; + } else { + return false; + } + if (index + length > text.size()) { + return false; + } + for (std::size_t offset = 1; offset < length; ++offset) { + const auto continuation = static_cast(text[index + offset]); + if ((continuation & 0xC0) != 0x80) { + return false; + } + codepoint = (codepoint << 6U) | (continuation & 0x3FU); + } + if ((length == 2 && codepoint < 0x80) || (length == 3 && codepoint < 0x800) || + (length == 4 && codepoint < 0x10000) || codepoint > 0x10FFFF || + (codepoint >= 0xD800 && codepoint <= 0xDFFF)) { + return false; + } + index += length; + } + return true; +} + +} // namespace + +// Peer-controlled text reaching a diagnostic message is a log-forging vector, and JSON is the sharp +// edge: a metadata document is decoded before its fields are interpolated, so an issuer written +// with `\r\n` escape sequences arrives as real control bytes. This one goes through +// MetadataPolicyError, whose constructor now flattens its own message so that no throw site has to +// remember to. +TEST(AuthDiagnosticsSanitizingTest, AnIssuerCarryingControlCharactersCannotForgeALogLine) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + const auto hostile_issuer = "https://evil" + forged_log_line() + "host.test"; + + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm.json") { + // `resource` identifies our server, so the document is accepted and the SDK goes on to + // the authorization server it names. That name is the payload. + return json_response({{"resource", base + "/mcp"}, + {"authorization_servers", json::array({hostile_issuer})}}); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(3), asio::detached); + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + std::promise result; + auto observed = result.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + std::string failure; + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + try { + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm.json")"); + failure = ""; + } catch (const std::exception& error) { + failure = error.what(); + } + result.set_value(failure); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + const auto failure = observed.get(); + ASSERT_NE(failure, ""); + // The control-character check is the assertion; it is what makes the forged second line + // impossible regardless of how the message is worded. + EXPECT_FALSE(carries_control_characters(failure)) << failure; + // The payload's visible text may still appear, flattened onto the SDK's own single line. What + // must not survive is its ability to start a line of its own. + EXPECT_EQ(failure.find('\n'), std::string::npos) << failure; + EXPECT_EQ(failure.find('\r'), std::string::npos) << failure; +} + +// The same property away from MetadataPolicyError, since sanitizing in that constructor covers only +// the throw sites that go through it. An `error` on the authorization response is peer-controlled +// too, and it is interpolated by validate_authorization_response() into the message run_challenge() +// throws. Being issuer-authentic makes it trustworthy as to origin, not as to content. +TEST(AuthDiagnosticsSanitizingTest, AnAuthorizationResponseErrorCannotForgeALogLine) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm.json") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(auth_server_metadata(base, true)); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(3), asio::detached); + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + // Echoes state and iss so the response is accepted as authentic, then reports the payload as + // the server's error. That ordering matters: the error is only read once the response has + // passed the state and iss checks. + auto hostile_callback = [](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + mcp::auth::AuthorizationResponse response; + response.state = request.state; + response.iss = request.issuer; + response.error = "access_denied" + forged_log_line() + "granted"; + co_return response; + }; + + std::promise result; + auto observed = result.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + std::string failure; + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + hostile_callback); + try { + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm.json")"); + failure = ""; + } catch (const std::exception& error) { + failure = error.what(); + } + result.set_value(failure); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + const auto failure = observed.get(); + ASSERT_NE(failure, ""); + EXPECT_NE(failure.find("Authorization response rejected"), std::string::npos) << failure; + EXPECT_FALSE(carries_control_characters(failure)) << failure; + EXPECT_EQ(failure.find('\n'), std::string::npos) << failure; + EXPECT_EQ(failure.find('\r'), std::string::npos) << failure; +} + +// The sanitizer's budget is counted in bytes, so a multi-byte character can straddle it; cutting +// there would emit invalid UTF-8, which a JSON log encoder throws on or drops. The payload is a long +// run of three-byte characters chosen so the 256-byte limit lands in the middle of one, plus +// ill-formed bytes, which reach the SDK through headers rather than through a parsed document. +TEST(AuthDiagnosticsSanitizingTest, ATruncatedMultiByteIssuerStaysWellFormedUtf8) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + // U+4E16 is three bytes in UTF-8. 200 of them is 600 bytes, comfortably past the budget, and + // 256 is not a multiple of 3, so the cut necessarily falls inside a character. + std::string wide; + for (int index = 0; index < 200; ++index) { + wide += "\xE4\xB8\x96"; + } + const auto hostile_issuer = "https://evil" + wide + ".test"; + + server.set_handler([&](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm.json") { + return json_response({{"resource", base + "/mcp"}, + {"authorization_servers", json::array({hostile_issuer})}}); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(3), asio::detached); + + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(server.origin()); + + std::promise result; + auto observed = result.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + std::string failure; + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + try { + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm.json")"); + failure = ""; + } catch (const std::exception& error) { + failure = error.what(); + } + result.set_value(failure); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + const auto failure = observed.get(); + ASSERT_NE(failure, ""); + // The property one level up from "no control characters": the message is something a log + // encoder can actually encode. + EXPECT_TRUE(is_well_formed_utf8(failure)) + << "sanitized diagnostic is not valid UTF-8, length " << failure.size(); + EXPECT_FALSE(carries_control_characters(failure)) << failure.size(); + // It was genuinely cut, so the boundary case was exercised rather than skipped. + EXPECT_NE(failure.find("..."), std::string::npos) << "payload did not reach the budget"; +} + +// The same guarantee for bytes that were never valid UTF-8 to begin with. A `Location` header is +// raw bytes with no parser between it and the SDK, unlike a JSON document, so this is the input +// class the truncation fix alone would not have covered. +TEST(AuthDiagnosticsSanitizingTest, IllFormedBytesAreReplacedRatherThanCopiedThrough) { + // A lone continuation byte, a truncated three-byte lead, and a surrogate encoding: each is + // ill-formed, and none is a control character, so the earlier assertions would all have passed. + // Built byte by byte rather than as literals: a hex escape in a C++ string literal swallows + // every following hex digit, so "\xC0after" is one out-of-range escape, not a byte and a word. + const std::string ill_formed = std::string(1, '\x80') + std::string(1, '\xE4') + + std::string(1, '\xB8') + std::string(1, '\xED') + + std::string(1, '\xA0') + std::string(1, '\x80'); + ASSERT_FALSE(is_well_formed_utf8(ill_formed)) << "the payload must really be ill-formed"; + + const auto cleaned = mcp::auth::detail::sanitize_for_diagnostics(ill_formed); + + EXPECT_TRUE(is_well_formed_utf8(cleaned)) << "ill-formed input survived into the diagnostic"; + EXPECT_FALSE(carries_control_characters(cleaned)); + // Valid text either side of the damage is preserved, so this is not simply dropping everything. + const auto mixed = + mcp::auth::detail::sanitize_for_diagnostics(std::string("before") + '\xC0' + "after"); + EXPECT_TRUE(is_well_formed_utf8(mixed)) << mixed; + EXPECT_NE(mixed.find("before"), std::string::npos) << mixed; + EXPECT_NE(mixed.find("after"), std::string::npos) << mixed; +} + +namespace { + +/// Encode one codepoint as UTF-8, so the bidi tests read as codepoints rather than as byte soup. +std::string utf8_codepoint(std::uint32_t codepoint) { + std::string encoded; + if (codepoint < 0x80) { + encoded.push_back(static_cast(codepoint)); + } else if (codepoint < 0x800) { + encoded.push_back(static_cast(0xC0U | (codepoint >> 6U))); + encoded.push_back(static_cast(0x80U | (codepoint & 0x3FU))); + } else if (codepoint < 0x10000) { + encoded.push_back(static_cast(0xE0U | (codepoint >> 12U))); + encoded.push_back(static_cast(0x80U | ((codepoint >> 6U) & 0x3FU))); + encoded.push_back(static_cast(0x80U | (codepoint & 0x3FU))); + } else { + encoded.push_back(static_cast(0xF0U | (codepoint >> 18U))); + encoded.push_back(static_cast(0x80U | ((codepoint >> 12U) & 0x3FU))); + encoded.push_back(static_cast(0x80U | ((codepoint >> 6U) & 0x3FU))); + encoded.push_back(static_cast(0x80U | (codepoint & 0x3FU))); + } + return encoded; +} + +} // namespace + +// Forgery without a line break. U+202E reverses the rendering of everything after it, so a peer +// that gets one into a refusal message can make the message read as its opposite in a terminal or +// a browser-based log viewer -- the same outcome U+2028/U+2029 were flattened to prevent, reached +// without ending a line at all. +TEST(AuthDiagnosticsSanitizingTest, BidirectionalOverridesAreFlattenedToo) { + // The bidi controls Unicode defines: the explicit embeddings and overrides (U+202A-U+202E), + // the isolates (U+2066-U+2069) and all three implicit marks -- LRM (U+200E), RLM (U+200F) and + // ALM (U+061C). ALM is the one an enumeration reaches for last, because it sits in the Arabic + // block rather than beside the other two; it is the same class as RLM. + const std::vector reordering = {0x202A, 0x202B, 0x202C, 0x202D, 0x202E, 0x2066, + 0x2067, 0x2068, 0x2069, 0x200E, 0x200F, 0x061C}; + + for (const auto codepoint : reordering) { + const auto raw = "server '" + utf8_codepoint(codepoint) + "denied'"; + const auto cleaned = mcp::auth::detail::sanitize_for_diagnostics(raw); + + EXPECT_TRUE(is_well_formed_utf8(cleaned)) << std::hex << codepoint; + EXPECT_EQ(cleaned.find(utf8_codepoint(codepoint)), std::string::npos) + << "U+" << std::hex << std::uppercase << codepoint + << " survived sanitization and can still re-order the rest of the line"; + // Flattened, not dropped: the surrounding text is still readable. + EXPECT_NE(cleaned.find("server '"), std::string::npos) << std::hex << codepoint; + EXPECT_NE(cleaned.find("denied'"), std::string::npos) << std::hex << codepoint; + } + + // Ordinary text that merely lives in the same planes is untouched: this is a targeted flatten, + // not a blanket refusal of non-ASCII. + const auto japanese = utf8_codepoint(0x65E5) + utf8_codepoint(0x672C); + EXPECT_EQ(mcp::auth::detail::sanitize_for_diagnostics(japanese), japanese); + const auto adjacent = utf8_codepoint(0x2029) + utf8_codepoint(0x202F) + utf8_codepoint(0x2065) + + utf8_codepoint(0x206A); + const auto cleaned_adjacent = mcp::auth::detail::sanitize_for_diagnostics(adjacent); + EXPECT_NE(cleaned_adjacent.find(utf8_codepoint(0x202F)), std::string::npos) + << "U+202F is a space, not an override, and must not be caught by an off-by-one range"; + EXPECT_NE(cleaned_adjacent.find(utf8_codepoint(0x2065)), std::string::npos) + << "U+2065 sits just below the isolates and must not be caught by an off-by-one range"; + EXPECT_NE(cleaned_adjacent.find(utf8_codepoint(0x206A)), std::string::npos) + << "U+206A sits just above the isolates and must not be caught by an off-by-one range"; + const auto arabic_neighbour = utf8_codepoint(0x061B); + EXPECT_EQ(mcp::auth::detail::sanitize_for_diagnostics(arabic_neighbour), arabic_neighbour) + << "U+061B is the Arabic semicolon, a printable character next to ALM, and must not be " + "caught by a range written around U+061C"; +} + +// --------------------------------------------------------------------------- +// Client against the real server. +// +// Every other test in this file drives LoopbackServer, whose answers were written alongside the +// client, so it can never disagree with the client. The tests below stand up the shipped +// StreamableHttpSessionManager and point the shipped OAuthAuthorizationManager at the challenge it +// really emits. +// --------------------------------------------------------------------------- + +namespace { + +/// The status and challenge of one unauthenticated request, read off the wire. +struct RawChallenge { + unsigned int status{0}; + bool had_www_authenticate{false}; + std::string www_authenticate; +}; + +mcp::Task fetch_unauthenticated_challenge(const asio::any_io_executor& executor, + unsigned short port) { + const auto body = json{{"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"clientInfo", {{"name", "auth-pairing-test"}, {"version", "1"}}}, + {"capabilities", json::object()}}}, + {"id", 1}} + .dump(); + + beast::tcp_stream stream(executor); + const std::vector endpoints{{asio::ip::make_address("127.0.0.1"), port}}; + co_await stream.async_connect(endpoints, asio::use_awaitable); + + http::request request(http::verb::post, "/mcp", + mcp::constants::g_http_version_11); + request.set(http::field::host, "127.0.0.1:" + std::to_string(port)); + request.set(http::field::content_type, "application/json"); + request.set(http::field::accept, "application/json, text/event-stream"); + request.set("MCP-Protocol-Version", std::string(mcp::g_LATEST_PROTOCOL_VERSION)); + request.body() = body; + request.prepare_payload(); + co_await http::async_write(stream, request, asio::use_awaitable); + + beast::flat_buffer buffer; + http::response response; + co_await http::async_read(stream, buffer, response, asio::use_awaitable); + + RawChallenge challenge; + challenge.status = response.result_int(); + const auto header = response.find(http::field::www_authenticate); + if (header != response.end()) { + challenge.had_www_authenticate = true; + challenge.www_authenticate = std::string(header->value()); + } + + beast::error_code ignored; + (void)stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, ignored); + co_return challenge; +} + +mcp::StreamableHttpSessionManager::ServerFactory make_pairing_server_factory() { + return [](const asio::any_io_executor&) -> std::unique_ptr { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + return std::make_unique(mcp::Implementation{"auth-pairing-server", "1.0.0"}, + std::move(caps)); + }; +} + +} // namespace + +// Pins how a server configured with only a validator behaves. RFC 9728 section 5.1 has the resource +// server point the client at its metadata with a `resource_metadata` parameter on the challenge; this +// server sends a bare `Bearer`, so the client can only fall back to the well-known location derived +// from its configured URL, which this server does not serve, and discovery fails. +// `ClientDiscoversAuthorizationFromTheServersOwnChallenge` below is the same pairing with +// set_protected_resource_metadata() configured. +TEST(AuthClientServerPairingTest, ServerChallengeCarriesNoResourceMetadataSoDiscoveryCannotStart) { + constexpr unsigned short port = 19211; + + asio::io_context io_ctx; + mcp::StreamableHttpSessionManager manager(io_ctx.get_executor(), "127.0.0.1", port, + make_pairing_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "valid-token"; }); + asio::co_spawn(io_ctx, manager.listen(), asio::detached); + + RawChallenge challenge; + bool client_authorized = false; + std::string client_failure; + std::atomic discovery_attempts{0}; + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + challenge = co_await fetch_unauthenticated_challenge(io_ctx.get_executor(), port); + + // The shipped client, handed the shipped server's own challenge. + const auto base = "http://127.0.0.1:" + std::to_string(port); + auto store = std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + config.server_url = base + "/mcp"; + config.client_id = "test-client"; + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.policy = loopback_policy(base); + // Called once per exchange, so it counts the discovery requests the client had to + // invent for itself. + config.host_resolver = [&](const std::string&, const std::string&) { + discovery_attempts.fetch_add(1); + return std::vector{"127.0.0.1"}; + }; + + mcp::auth::OAuthAuthorizationManager authorization(io_ctx.get_executor(), store, config, + echoing_callback(nullptr)); + try { + client_authorized = + co_await authorization.try_handle_challenge(challenge.www_authenticate); + } catch (const std::exception& error) { + client_failure = error.what(); + } + manager.close(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_EQ(challenge.status, 401U); + ASSERT_TRUE(challenge.had_www_authenticate); + // The challenge is bare. + EXPECT_EQ(challenge.www_authenticate, "Bearer") + << "an unconfigured server must not advertise challenge parameters"; + EXPECT_EQ(challenge.www_authenticate.find("resource_metadata"), std::string::npos) + << challenge.www_authenticate; + + // Told nothing, the client does not give up: it falls back to guessing the well-known locations + // under the URL it was configured with, spends real requests on them, and only then fails. So the + // cost of the missing parameter is not one failed handshake, it is the client probing a server + // that never advertised anything. + EXPECT_GT(discovery_attempts.load(), 0) + << "the client should have been driven to guess at well-known locations"; + EXPECT_FALSE(client_authorized) + << "the client authorized against a challenge that names no metadata location"; + // It reports the configured server URL, because that is all it ever had to go on; a compliant + // challenge would have named the document location instead. + EXPECT_NE(client_failure.find("Failed to discover protected resource metadata"), std::string::npos) + << client_failure; + EXPECT_NE(client_failure.find("/mcp"), std::string::npos) << client_failure; +} + +// A server that is told what resource it represents advertises where its metadata lives, so the +// shipped client has somewhere to begin discovery. +// +// The seam stays guarded in both directions: the test above pins that a server configured with only a +// validator sends the bare challenge, so the parameter cannot appear by accident, and this one fails +// if it ever stops appearing when the metadata IS configured. +TEST(AuthClientServerPairingTest, ClientDiscoversAuthorizationFromTheServersOwnChallenge) { + constexpr unsigned short port = 19212; + + asio::io_context io_ctx; + mcp::StreamableHttpSessionManager manager(io_ctx.get_executor(), "127.0.0.1", port, + make_pairing_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "valid-token"; }); + + // What the resource is called from outside is stated, never inferred from the bound address. + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "http://127.0.0.1:" + std::to_string(port) + "/mcp"; + metadata.authorization_servers = {"https://auth.example.com"}; + manager.set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx, manager.listen(), asio::detached); + + RawChallenge challenge; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + challenge = co_await fetch_unauthenticated_challenge(io_ctx.get_executor(), port); + manager.close(); + }, + asio::detached); + + io_ctx.run(); + + ASSERT_EQ(challenge.status, 401U); + ASSERT_TRUE(challenge.had_www_authenticate); + const auto parsed = mcp::auth::parse_www_authenticate(challenge.www_authenticate); + const auto bearer = mcp::auth::select_bearer_challenge(parsed); + ASSERT_TRUE(bearer.has_value()); + EXPECT_TRUE(bearer->resource_metadata.has_value()) + << "RFC 9728 5.1: the challenge must name where the client can discover how to authenticate"; +} + +// set_metadata_policy() and set_host_resolver() are taken under the mutex that guards the exchange +// list, and each exchange reads them once when it is built. An exchange re-validates every redirect +// hop, so this pins that a chain runs under one policy from end to end even when the policy is +// swapped mid-chain. +TEST(AuthHttpClientPolicyTest, APolicyInstalledMidExchangeDoesNotChangeTheRulesUnderIt) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/hop0") { + http::response redirect(http::status::found, request.version()); + redirect.set(http::field::location, base + "/hop1"); + return redirect; + } + return json_response({{"arrived", true}}); + }); + asio::co_spawn(io_ctx, server.serve(4), asio::detached); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(server.origin())); + + // The resolver hook runs synchronously inside the exchange, which makes "part-way through" an + // exact point rather than a race. On the first hop it swaps in a policy that allows nothing. + std::atomic resolver_calls{0}; + auto* raw_client = client.get(); + client->set_host_resolver([&resolver_calls, raw_client](const std::string&, const std::string&) { + if (resolver_calls.fetch_add(1) == 0) { + raw_client->set_metadata_policy(mcp::auth::MetadataFetchPolicy{}); + } + return std::vector{"127.0.0.1"}; + }); + + std::promise result; + auto observed = result.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + std::string outcome; + try { + const auto body = co_await client->get_json(base + "/hop0"); + outcome = body.value("arrived", false) ? "arrived" : "unexpected body"; + } catch (const std::exception& error) { + outcome = std::string("threw: ") + error.what(); + } + result.set_value(outcome); + server.close(); + }, + asio::detached); + + io_ctx.run(); + + EXPECT_EQ(observed.get(), "arrived") + << "the redirect chain was re-validated against a policy installed after it started"; + const std::vector expected_targets{"/hop0", "/hop1"}; + EXPECT_EQ(server.targets(), expected_targets); + EXPECT_GE(resolver_calls.load(), 2) << "both hops must have resolved for this to prove anything"; +} + +// The race itself. Its value is under ThreadSanitizer, where the unsynchronised version reports the +// setters against the reads in run_get() and connect(); without a sanitizer it still asserts that +// nothing hangs or crashes while a policy and a resolver are replaced under live requests. +TEST(AuthHttpClientPolicyTest, ReplacingThePolicyAndResolverUnderLiveRequestsIsSafe) { + constexpr int io_thread_count = 4; + constexpr int request_count = 60; + + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + server.set_handler( + [](const http::request&) { return json_response({{"ok", true}}); }); + asio::co_spawn(io_ctx, server.serve(request_count + 4), asio::detached); + + auto client = std::make_shared(io_ctx.get_executor()); + client->set_metadata_policy(loopback_policy(server.origin())); + + std::atomic completed{0}; + std::atomic stop_writing{false}; + + std::vector runners; + runners.reserve(io_thread_count); + for (int index = 0; index < io_thread_count; ++index) { + runners.emplace_back([&io_ctx]() { io_ctx.run(); }); + } + + // An application thread that keeps installing both, exactly as a direct user of this client + // might when its configuration changes. Every policy it installs is equivalent, so a request + // succeeds whichever one it pinned; what is under test is the concurrent write, not the outcome. + std::thread writer([&]() { + while (!stop_writing.load()) { + client->set_metadata_policy(loopback_policy(server.origin())); + client->set_host_resolver(nullptr); + } + }); + + for (int index = 0; index < request_count; ++index) { + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + (void)co_await client->get_json(base + "/probe"); + } catch (...) { + // A refusal is an acceptable outcome; a race is not, and that is TSan's call. + } + completed.fetch_add(1); + }, + asio::detached); + } + + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(30); + while (completed.load() < request_count && std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + stop_writing.store(true); + writer.join(); + server.close(); + io_ctx.stop(); + for (auto& runner : runners) { + runner.join(); + } + + EXPECT_EQ(completed.load(), request_count) << "a request neither completed nor failed"; +} + +// A protected-resource document names its authorization servers, and nothing validates that a name +// is a URL before it reaches a diagnostic. These two tests hold the line at the exception message: +// whatever the peer put in `authorization_servers`, what surfaces to the application must be one +// line and must be bounded. +namespace { + +/// Drive `try_handle_challenge` against a `/prm` document carrying `authorization_server` verbatim, +/// and return `what()` from whatever it throws. An empty string means it did not throw. +std::string challenge_failure_message(const std::string& authorization_server) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base, &authorization_server](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response({{"resource", base + "/mcp"}, + {"authorization_servers", json::array({authorization_server})}}); + } + return status_response(http::status::not_found); + }); + asio::co_spawn(io_ctx, server.serve(1), asio::detached); + + ManagerFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.client_id = "test-client"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.policy = loopback_policy(server.origin()); + + std::string message; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.store, + fixture.config, echoing_callback(nullptr)); + try { + (void)co_await manager.try_handle_challenge( + R"(Bearer realm="mcp", resource_metadata=")" + base + R"(/prm")"); + } catch (const std::exception& error) { + message = error.what(); + } + server.close(); + }, + asio::detached); + + io_ctx.run(); + return message; +} + +} // namespace + +TEST(AuthDiagnosticSanitizationTest, DoesNotLeaveLineBreaksFromTheAuthorizationServerInDiagnostics) { + // No "://" anywhere in the payload: the identifier must reach the scheme check as the thing it + // is, a bare string. A payload carrying a scheme parses instead and is refused by the metadata + // policy, whose own message is sanitized, which would make this test pass without proving + // anything. + const auto message = challenge_failure_message("evil\r\nFORGED"); + + ASSERT_FALSE(message.empty()) << "the malformed authorization server identifier was accepted"; + ASSERT_NE(message.find("FORGED"), std::string::npos) + << "the identifier never reached the diagnostic, so this test proves nothing: " << message; + EXPECT_EQ(message.find('\r'), std::string::npos) + << "a carriage return the peer chose reached the diagnostic: " << message; + EXPECT_EQ(message.find('\n'), std::string::npos) + << "a line feed the peer chose reached the diagnostic: " << message; +} + +TEST(AuthDiagnosticSanitizationTest, CapsTheAuthorizationServerIdentifierItReportsInDiagnostics) { + constexpr std::size_t oversized_length = 64U * 1024U; + const auto message = challenge_failure_message(std::string(oversized_length, 'A')); + + ASSERT_FALSE(message.empty()) << "the malformed authorization server identifier was accepted"; + // The sanitizer's bound is 256 bytes plus an ellipsis; the surrounding literal is short, so any + // message near the peer's own length means the value arrived unbounded. + EXPECT_LT(message.size(), 1024U) + << "the peer's " << oversized_length << "-byte identifier was not capped; message is " + << message.size() << " bytes"; +} diff --git a/test/auth/auth_challenge_test.cpp b/test/auth/auth_challenge_test.cpp new file mode 100644 index 0000000..dbebfd6 --- /dev/null +++ b/test/auth/auth_challenge_test.cpp @@ -0,0 +1,642 @@ +/** + * @file auth_challenge_test.cpp + * @brief Tests for WWW-Authenticate challenge parsing, RFC 9207 issuer validation, and the + * outbound-request policy that guards OAuth metadata discovery. + */ + +#include + +#include +#include +#include +#include + +namespace { + +mcp::auth::AuthorizationRequest make_request(std::string issuer, bool issuer_parameter_supported) { + mcp::auth::AuthorizationRequest request; + request.state = "recorded-state"; + request.code_verifier = "recorded-verifier"; + request.issuer = std::move(issuer); + request.issuer_parameter_supported = issuer_parameter_supported; + return request; +} + +mcp::auth::AuthorizationResponse make_response(std::optional iss) { + mcp::auth::AuthorizationResponse response; + response.code = "auth-code"; + response.state = "recorded-state"; + response.iss = std::move(iss); + return response; +} + +mcp::auth::MetadataFetchPolicy allow_origin(std::string origin) { + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins.push_back(std::move(origin)); + return policy; +} + +} // namespace + +TEST(AuthChallengeParserTest, ParsesSchemeWithNoParameters) { + const auto challenges = mcp::auth::parse_www_authenticate("Bearer"); + ASSERT_EQ(challenges.size(), 1U); + EXPECT_EQ(challenges[0].scheme, "Bearer"); + EXPECT_TRUE(challenges[0].is_bearer()); + EXPECT_TRUE(challenges[0].parameters.empty()); +} + +TEST(AuthChallengeParserTest, ParsesQuotedAndUnquotedValues) { + const auto challenges = + mcp::auth::parse_www_authenticate(R"(Bearer realm="mcp server", error=invalid_token)"); + ASSERT_EQ(challenges.size(), 1U); + EXPECT_EQ(challenges[0].realm, "mcp server"); + EXPECT_EQ(challenges[0].error, "invalid_token"); +} + +TEST(AuthChallengeParserTest, ExpandsBackslashEscapesInQuotedValues) { + const auto challenges = mcp::auth::parse_www_authenticate( + R"(Bearer error_description="the server said \"no\" \\ twice")"); + ASSERT_EQ(challenges.size(), 1U); + EXPECT_EQ(challenges[0].error_description, R"(the server said "no" \ twice)"); +} + +TEST(AuthChallengeParserTest, MatchesParameterNamesCaseInsensitively) { + const auto challenges = mcp::auth::parse_www_authenticate( + R"(Bearer REALM="r", Resource_Metadata="https://rs.test/.well-known/x", SCOPE="a b")"); + ASSERT_EQ(challenges.size(), 1U); + EXPECT_EQ(challenges[0].realm, "r"); + EXPECT_EQ(challenges[0].resource_metadata, "https://rs.test/.well-known/x"); + EXPECT_EQ(challenges[0].scope, "a b"); + ASSERT_EQ(challenges[0].parameters.size(), 3U); + EXPECT_EQ(challenges[0].parameters[0].first, "realm"); + EXPECT_EQ(challenges[0].parameters[1].first, "resource_metadata"); + EXPECT_EQ(challenges[0].parameters[2].first, "scope"); +} + +TEST(AuthChallengeParserTest, AcceptsArbitraryParameterOrdering) { + const auto forward = mcp::auth::parse_www_authenticate( + R"(Bearer resource_metadata="https://rs.test/prm", scope="a", realm="r")"); + const auto reverse = mcp::auth::parse_www_authenticate( + R"(Bearer realm="r", scope="a", resource_metadata="https://rs.test/prm")"); + ASSERT_EQ(forward.size(), 1U); + ASSERT_EQ(reverse.size(), 1U); + EXPECT_EQ(forward[0].resource_metadata, reverse[0].resource_metadata); + EXPECT_EQ(forward[0].scope, reverse[0].scope); + EXPECT_EQ(forward[0].realm, reverse[0].realm); +} + +TEST(AuthChallengeParserTest, SplitsSeveralChallengesInOneHeader) { + const auto challenges = + mcp::auth::parse_www_authenticate(R"(Basic realm="legacy", Bearer realm="mcp")"); + ASSERT_EQ(challenges.size(), 2U); + EXPECT_EQ(challenges[0].scheme, "Basic"); + EXPECT_EQ(challenges[0].realm, "legacy"); + EXPECT_FALSE(challenges[0].is_bearer()); + EXPECT_EQ(challenges[1].scheme, "Bearer"); + EXPECT_EQ(challenges[1].realm, "mcp"); +} + +TEST(AuthChallengeParserTest, SkipsToken68CredentialsBeforeTheNextChallenge) { + const auto challenges = + mcp::auth::parse_www_authenticate(R"(Negotiate a1b2c3==, Bearer realm="mcp")"); + ASSERT_EQ(challenges.size(), 2U); + EXPECT_EQ(challenges[0].scheme, "Negotiate"); + EXPECT_EQ(challenges[1].scheme, "Bearer"); + EXPECT_EQ(challenges[1].realm, "mcp"); +} + +TEST(AuthChallengeParserTest, CombinesSeveralHeaderValues) { + const std::vector headers = {R"(Basic realm="legacy")", + R"(Bearer resource_metadata="https://rs.test/prm")"}; + const auto challenges = mcp::auth::parse_www_authenticate(headers); + ASSERT_EQ(challenges.size(), 2U); + EXPECT_EQ(challenges[0].scheme, "Basic"); + EXPECT_EQ(challenges[1].resource_metadata, "https://rs.test/prm"); +} + +TEST(AuthChallengeParserTest, KeepsFirstOccurrenceOfARepeatedParameter) { + const auto challenges = + mcp::auth::parse_www_authenticate(R"(Bearer realm="first", realm="second")"); + ASSERT_EQ(challenges.size(), 1U); + EXPECT_EQ(challenges[0].realm, "first"); + EXPECT_EQ(challenges[0].parameters.size(), 2U); +} + +TEST(AuthChallengeParserTest, DiscardsTrailingDamageAfterAnUnterminatedQuotedString) { + const auto challenges = mcp::auth::parse_www_authenticate(R"(Bearer realm="unterminated)"); + ASSERT_EQ(challenges.size(), 1U); + EXPECT_EQ(challenges[0].scheme, "Bearer"); + EXPECT_FALSE(challenges[0].realm.has_value()); +} + +TEST(AuthChallengeParserTest, SelectsTheFirstBearerChallenge) { + const auto challenges = mcp::auth::parse_www_authenticate( + R"(Basic realm="legacy", Bearer realm="one", Bearer realm="two")"); + const auto selected = mcp::auth::select_bearer_challenge(challenges); + ASSERT_TRUE(selected.has_value()); + EXPECT_EQ(selected->realm, "one"); +} + +TEST(AuthChallengeParserTest, SelectsNothingWhenNoBearerChallengeIsPresent) { + const auto challenges = mcp::auth::parse_www_authenticate(R"(Basic realm="legacy")"); + EXPECT_FALSE(mcp::auth::select_bearer_challenge(challenges).has_value()); +} + +TEST(AuthAuthorizationResponseParseTest, DecodesRedirectQueryParameters) { + const auto response = mcp::auth::parse_authorization_response( + "http://127.0.0.1:9000/callback?code=abc%2F123&state=s1&iss=https%3A%2F%2Fas.test"); + EXPECT_EQ(response.code, "abc/123"); + EXPECT_EQ(response.state, "s1"); + EXPECT_EQ(response.iss, "https://as.test"); +} + +TEST(AuthAuthorizationResponseParseTest, DecodesPlusAsSpaceAndBareQueryStrings) { + const auto response = + mcp::auth::parse_authorization_response("error=access_denied&error_description=user+said+no"); + EXPECT_EQ(response.error, "access_denied"); + EXPECT_EQ(response.error_description, "user said no"); + EXPECT_FALSE(response.code.has_value()); +} + +TEST(AuthAuthorizationResponseParseTest, IgnoresTheFragmentComponent) { + const auto response = + mcp::auth::parse_authorization_response("http://127.0.0.1/cb?code=abc#state=spoofed"); + EXPECT_EQ(response.code, "abc"); + EXPECT_FALSE(response.state.has_value()); +} + +// The RFC 9207 Section 2.4 decision table as adopted by MCP: four rows over +// `authorization_response_iss_parameter_supported` by the presence of `iss`. +TEST(AuthIssuerValidationTest, AppliesTheFourRowDecisionTable) { + struct Row { + const char* name; + bool issuer_parameter_supported; + bool iss_present; + mcp::auth::AuthorizationResponseStatus expected; + }; + const Row rows[] = { + {"advertised and present compares equal", true, true, + mcp::auth::AuthorizationResponseStatus::accepted}, + {"advertised and absent rejects", true, false, + mcp::auth::AuthorizationResponseStatus::issuer_missing}, + {"not advertised but present still compares", false, true, + mcp::auth::AuthorizationResponseStatus::accepted}, + {"not advertised and absent proceeds", false, false, + mcp::auth::AuthorizationResponseStatus::accepted}, + }; + + for (const auto& row : rows) { + const auto request = make_request("https://as.test/tenant1", row.issuer_parameter_supported); + const auto response = make_response( + row.iss_present ? std::optional("https://as.test/tenant1") : std::nullopt); + const auto validation = mcp::auth::validate_authorization_response(request, response); + EXPECT_EQ(validation.status, row.expected) << row.name; + } +} + +TEST(AuthIssuerValidationTest, RejectsAPresentIssuerThatDiffers) { + for (const bool advertised : {true, false}) { + const auto request = make_request("https://as.test/tenant1", advertised); + const auto response = make_response("https://evil.test/tenant1"); + const auto validation = mcp::auth::validate_authorization_response(request, response); + EXPECT_EQ(validation.status, mcp::auth::AuthorizationResponseStatus::issuer_mismatch) + << "advertised=" << advertised; + } +} + +// Comparison is plain string equality. Each pair below is equivalent under RFC 3986 syntax-based +// normalization and must still be rejected. +TEST(AuthIssuerValidationTest, ComparisonIsNotCanonicalizing) { + struct Pair { + const char* name; + const char* recorded; + const char* returned; + }; + const Pair pairs[] = { + {"trailing slash", "https://as.test/tenant1", "https://as.test/tenant1/"}, + {"host case folding", "https://as.test/tenant1", "https://AS.test/tenant1"}, + {"scheme case folding", "https://as.test/tenant1", "HTTPS://as.test/tenant1"}, + {"default port elision", "https://as.test/tenant1", "https://as.test:443/tenant1"}, + {"percent-encoding", "https://as.test/tenant~1", "https://as.test/tenant%7E1"}, + }; + + for (const auto& pair : pairs) { + const auto request = make_request(pair.recorded, true); + const auto response = make_response(pair.returned); + const auto validation = mcp::auth::validate_authorization_response(request, response); + EXPECT_EQ(validation.status, mcp::auth::AuthorizationResponseStatus::issuer_mismatch) + << pair.name; + } +} + +TEST(AuthIssuerValidationTest, RejectsErrorResponsesBeforeActingOnTheirErrorValues) { + auto response = make_response("https://evil.test"); + response.code.reset(); + response.error = "access_denied"; + response.error_description = "attacker supplied text"; + response.error_uri = "https://evil.test/explain"; + + const auto validation = + mcp::auth::validate_authorization_response(make_request("https://as.test", true), response); + EXPECT_EQ(validation.status, mcp::auth::AuthorizationResponseStatus::issuer_mismatch); + EXPECT_EQ(validation.message.find("access_denied"), std::string::npos); + EXPECT_EQ(validation.message.find("attacker supplied text"), std::string::npos); +} + +TEST(AuthIssuerValidationTest, ReportsAServerErrorOnlyWhenTheIssuerIsAuthentic) { + auto response = make_response("https://as.test"); + response.code.reset(); + response.error = "access_denied"; + + const auto validation = + mcp::auth::validate_authorization_response(make_request("https://as.test", true), response); + EXPECT_EQ(validation.status, mcp::auth::AuthorizationResponseStatus::server_error); +} + +TEST(AuthIssuerValidationTest, RejectsMissingAndMismatchedState) { + auto missing = make_response("https://as.test"); + missing.state.reset(); + EXPECT_EQ(mcp::auth::validate_authorization_response(make_request("https://as.test", true), missing) + .status, + mcp::auth::AuthorizationResponseStatus::state_missing); + + auto mismatched = make_response("https://as.test"); + mismatched.state = "other-state"; + EXPECT_EQ( + mcp::auth::validate_authorization_response(make_request("https://as.test", true), mismatched) + .status, + mcp::auth::AuthorizationResponseStatus::state_mismatch); +} + +TEST(AuthIssuerValidationTest, RejectsAnAuthenticResponseWithNoCode) { + auto response = make_response("https://as.test"); + response.code.reset(); + EXPECT_EQ( + mcp::auth::validate_authorization_response(make_request("https://as.test", true), response) + .status, + mcp::auth::AuthorizationResponseStatus::code_missing); +} + +TEST(AuthMetadataPolicyTest, DefaultPolicyDeniesEveryOrigin) { + const mcp::auth::MetadataFetchPolicy policy; + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test/.well-known/x"), + mcp::auth::MetadataUrlDecision::origin_not_allowed); +} + +TEST(AuthMetadataPolicyTest, DenyListWinsOverAllowList) { + auto policy = allow_origin("https://as.test"); + policy.denied_origins.emplace_back("https://as.test"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, OriginComparisonIgnoresExplicitDefaultPort) { + const auto policy = allow_origin("https://as.test"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test/prm"), + mcp::auth::MetadataUrlDecision::allowed); + // An explicit default port canonicalizes to the same origin as no port at all. + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test:443/prm"), + mcp::auth::MetadataUrlDecision::allowed); + // A subdomain is a genuinely different origin and must stay refused. + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://sub.as.test/prm"), + mcp::auth::MetadataUrlDecision::origin_not_allowed); + // A non-default port is preserved and still distinguishes origins. + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test:8443/prm"), + mcp::auth::MetadataUrlDecision::origin_not_allowed); +} + +TEST(AuthMetadataPolicyTest, DenyListCanonicalizesSchemeAndHostCase) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "HTTPS://Evil.example/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, DenyListCanonicalizesExplicitDefaultPort) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://evil.example:443/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, DenyListCanonicalizesTrailingDot) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://evil.example./prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, AllowListEntryCaseDoesNotHaveToMatchTheUrl) { + const auto policy = allow_origin("HTTPS://AS.TEST"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test/prm"), + mcp::auth::MetadataUrlDecision::allowed); +} + +TEST(AuthMetadataPolicyTest, OriginAllowanceCallbackReceivesTheCanonicalOrigin) { + mcp::auth::MetadataFetchPolicy policy; + std::string observed; + policy.origin_allowance = [&observed](const std::string& origin) { + observed = origin; + return true; + }; + // The scheme is kept lowercase here because the separate https-required check further down is + // deliberately byte-exact; only the host case, the explicit default port and the origin passed to + // the callback are under test. + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://Evil.example:443/prm"), + mcp::auth::MetadataUrlDecision::allowed); + EXPECT_EQ(observed, "https://evil.example"); +} + +TEST(AuthMetadataPolicyTest, IpLiteralAndIpv6OriginsAreUnaffectedByCanonicalization) { + // An IP literal has no FQDN trailing dot to strip, and dotted decimal notation is not affected + // by lowercasing. + const auto v4_policy = allow_origin("https://93.184.216.34"); + EXPECT_EQ(mcp::auth::validate_metadata_url(v4_policy, "https://93.184.216.34/prm"), + mcp::auth::MetadataUrlDecision::allowed); + + // An IPv6 bracketed literal keeps its brackets; only the hex digits are lowercased. + auto v6_policy = allow_origin("https://[2606:2800:220:1::1]"); + EXPECT_EQ(mcp::auth::validate_metadata_url(v6_policy, "https://[2606:2800:220:1::1]/prm"), + mcp::auth::MetadataUrlDecision::allowed); + EXPECT_EQ(mcp::auth::validate_metadata_url(v6_policy, "https://[2606:2800:220:1::1]:443/prm"), + mcp::auth::MetadataUrlDecision::allowed); + + v6_policy.denied_origins.emplace_back("https://[2606:2800:220:1::1]"); + EXPECT_EQ(mcp::auth::validate_metadata_url(v6_policy, "https://[2606:2800:220:1::1]/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, DenyListCanonicalizesLeadingZerosInADefaultPort) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://evil.example:00443/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://evil.example:0443/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, NonDefaultPortCanonicalizesNumericallyDespiteLeadingZeros) { + const auto policy = allow_origin("https://evil.example:8443"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://evil.example:08443/prm"), + mcp::auth::MetadataUrlDecision::allowed); +} + +TEST(AuthMetadataPolicyTest, RejectsAPortThatIsNotAPlainDecimalNumber) { + const auto policy = allow_origin("https://evil.example"); + const char* urls[] = { + "https://evil.example:+443/prm", + "https://evil.example: 443/prm", + "https://evil.example:44a3/prm", + "https://evil.example:/prm", + }; + for (const auto* url : urls) { + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, url), + mcp::auth::MetadataUrlDecision::malformed_url) + << url; + } +} + +TEST(AuthMetadataPolicyTest, DenyListMatchesAnAlternateIpv6TextualForm) { + mcp::auth::MetadataFetchPolicy policy; + policy.denied_origins.emplace_back("https://[2606:2800:220:1::1]"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, + "https://[2606:2800:0220:0001:0000:0000:0000:0001]/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, AllowListEntryWithAPathMatchesNothing) { + const auto policy = allow_origin("https://as.test/realms/foo"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test/prm"), + mcp::auth::MetadataUrlDecision::origin_not_allowed); +} + +// An allow entry that is not a bare origin grants nothing, and says so by returning a decision +// rather than throwing: dropping it is fail-closed, so the quiet refusal is the safe outcome and +// stays the documented one. The deny list, where dropping an entry would be fail-open, is loud +// instead; see the deny-list tests below. +TEST(AuthMetadataPolicyTest, AllowListEntryThatIsNotABareOriginStillFailsClosedQuietly) { + const char* entries[] = { + "https://as.test/", "https://as.test/realms/foo", "https://as.test?tenant=1", + "https://as.test#frag", "https://as.test:44a3", + }; + for (const auto* entry : entries) { + const auto policy = allow_origin(entry); + mcp::auth::MetadataUrlDecision decision{}; + EXPECT_NO_THROW(decision = mcp::auth::validate_metadata_url(policy, "https://as.test/prm")) + << entry; + EXPECT_EQ(decision, mcp::auth::MetadataUrlDecision::origin_not_allowed) << entry; + } +} + +// A deny entry written with a trailing slash must not be discarded in silence, which would admit +// the origin the author meant to block. It is refused as a policy error instead. +TEST(AuthMetadataPolicyTest, DenyListEntryWithATrailingSlashIsRejected) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example/"); + try { + const auto decision = mcp::auth::validate_metadata_url(policy, "https://evil.example/prm"); + FAIL() << "expected MetadataPolicyError, got decision " << static_cast(decision); + } catch (const mcp::auth::MetadataPolicyError& error) { + EXPECT_EQ(error.decision(), mcp::auth::MetadataUrlDecision::denied_origin_entry_malformed); + EXPECT_EQ(error.target(), "https://evil.example/"); + } +} + +TEST(AuthMetadataPolicyTest, DenyListEntryWithAPathQueryOrFragmentIsRejected) { + const char* entries[] = { + "https://evil.example/realms/foo", "https://evil.example/?tenant=1", + "https://evil.example?tenant=1", "https://evil.example#frag", + "https://evil.example:44a3", + }; + for (const auto* entry : entries) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back(entry); + try { + const auto decision = mcp::auth::validate_metadata_url(policy, "https://evil.example/prm"); + ADD_FAILURE() << entry << " was admitted with decision " << static_cast(decision); + } catch (const mcp::auth::MetadataPolicyError& error) { + EXPECT_EQ(error.decision(), mcp::auth::MetadataUrlDecision::denied_origin_entry_malformed) + << entry; + EXPECT_EQ(error.target(), entry); + } + } +} + +// The refusal is a property of the policy, not of the URL: a malformed deny entry refuses every +// target, including one that no entry in the list was ever meant to describe. +TEST(AuthMetadataPolicyTest, MalformedDenyEntryRefusesAnUnrelatedTargetToo) { + auto policy = allow_origin("https://as.test"); + policy.denied_origins.emplace_back("https://evil.example/"); + EXPECT_THROW((void)mcp::auth::validate_metadata_url(policy, "https://as.test/prm"), + mcp::auth::MetadataPolicyError); +} + +// A malformed entry is reported wherever it sits in the list, even behind an entry that already +// matched, so a single bad rule cannot hide behind a working one. +TEST(AuthMetadataPolicyTest, MalformedDenyEntryIsReportedEvenAfterAMatchingEntry) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example"); + policy.denied_origins.emplace_back("https://other.example/"); + try { + (void)mcp::auth::validate_metadata_url(policy, "https://evil.example/prm"); + FAIL() << "expected MetadataPolicyError"; + } catch (const mcp::auth::MetadataPolicyError& error) { + EXPECT_EQ(error.decision(), mcp::auth::MetadataUrlDecision::denied_origin_entry_malformed); + EXPECT_EQ(error.target(), "https://other.example/"); + } +} + +// A malformed entry must not turn every deny list into an error: a well-formed list still denies +// exactly what it names, and still admits everything else the allow list permits. +TEST(AuthMetadataPolicyTest, WellFormedDenyListStillDeniesAndStillAdmits) { + auto policy = allow_origin("https://as.test"); + policy.allowed_origins.emplace_back("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example"); + policy.denied_origins.emplace_back("https://[2606:2800:220:1::1]:8443"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://evil.example/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test/prm"), + mcp::auth::MetadataUrlDecision::allowed); +} + +// A deny list the caller never populated is the common case and must stay free of policy errors. +TEST(AuthMetadataPolicyTest, EmptyDenyListIsNotAPolicyError) { + const auto policy = allow_origin("https://as.test"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test/prm"), + mcp::auth::MetadataUrlDecision::allowed); +} + +TEST(AuthMetadataPolicyTest, LoopbackOptOutIsCaseInsensitiveOnSchemeAndHost) { + auto policy = allow_origin("http://localhost:9000"); + policy.allow_plain_http_loopback = true; + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "HTTP://localhost:9000/prm"), + mcp::auth::MetadataUrlDecision::allowed); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://LOCALHOST:9000/prm"), + mcp::auth::MetadataUrlDecision::allowed); +} + +TEST(AuthMetadataPolicyTest, DenyListCanonicalizesMultipleTrailingDots) { + auto policy = allow_origin("https://evil.example"); + policy.denied_origins.emplace_back("https://evil.example"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://evil.example../prm"), + mcp::auth::MetadataUrlDecision::origin_denied); +} + +TEST(AuthMetadataPolicyTest, RejectsAHostThatIsNothingButDots) { + const mcp::auth::MetadataFetchPolicy policy; + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://.../prm"), + mcp::auth::MetadataUrlDecision::malformed_url); +} + +TEST(AuthMetadataPolicyTest, RejectsPlainHttpForNonLoopbackHosts) { + const auto policy = allow_origin("http://as.test"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://as.test/prm"), + mcp::auth::MetadataUrlDecision::scheme_not_allowed); +} + +TEST(AuthMetadataPolicyTest, LoopbackOptOutIsNeverImplicit) { + auto policy = allow_origin("http://127.0.0.1:9000"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://127.0.0.1:9000/prm"), + mcp::auth::MetadataUrlDecision::scheme_not_allowed); + + policy.allow_plain_http_loopback = true; + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://127.0.0.1:9000/prm"), + mcp::auth::MetadataUrlDecision::allowed); +} + +TEST(AuthMetadataPolicyTest, LoopbackOptOutDoesNotRelaxNonLoopbackHosts) { + auto policy = allow_origin("http://as.test"); + policy.allow_plain_http_loopback = true; + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://as.test/prm"), + mcp::auth::MetadataUrlDecision::scheme_not_allowed); +} + +TEST(AuthMetadataPolicyTest, RejectsUserinfoInTheAuthority) { + auto policy = allow_origin("https://as.test"); + policy.allowed_origins.emplace_back("https://as.test@169.254.169.254"); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://as.test@169.254.169.254/prm"), + mcp::auth::MetadataUrlDecision::malformed_url); +} + +TEST(AuthMetadataPolicyTest, RejectsBlockedAddressesWrittenAsUrlLiterals) { + struct Case { + const char* url; + const char* origin; + mcp::auth::MetadataUrlDecision expected; + }; + const Case cases[] = { + {"https://169.254.169.254/latest/meta-data", "https://169.254.169.254", + mcp::auth::MetadataUrlDecision::address_link_local}, + {"https://10.0.0.1/prm", "https://10.0.0.1", mcp::auth::MetadataUrlDecision::address_private}, + {"https://172.16.0.1/prm", "https://172.16.0.1", + mcp::auth::MetadataUrlDecision::address_private}, + {"https://192.168.1.1/prm", "https://192.168.1.1", + mcp::auth::MetadataUrlDecision::address_private}, + {"https://127.0.0.1/prm", "https://127.0.0.1", + mcp::auth::MetadataUrlDecision::address_loopback}, + {"https://224.0.0.1/prm", "https://224.0.0.1", + mcp::auth::MetadataUrlDecision::address_multicast}, + {"https://0.0.0.0/prm", "https://0.0.0.0", mcp::auth::MetadataUrlDecision::address_reserved}, + {"https://100.64.0.1/prm", "https://100.64.0.1", + mcp::auth::MetadataUrlDecision::address_reserved}, + {"https://[fe80::1]/prm", "https://[fe80::1]", + mcp::auth::MetadataUrlDecision::address_link_local}, + {"https://[fc00::1]/prm", "https://[fc00::1]", mcp::auth::MetadataUrlDecision::address_private}, + {"https://[::1]/prm", "https://[::1]", mcp::auth::MetadataUrlDecision::address_loopback}, + }; + + for (const auto& item : cases) { + EXPECT_EQ(mcp::auth::validate_metadata_url(allow_origin(item.origin), item.url), item.expected) + << item.url; + } +} + +TEST(AuthMetadataPolicyTest, ClassifiesIpv4MappedIpv6Addresses) { + const mcp::auth::MetadataFetchPolicy policy; + EXPECT_EQ(mcp::auth::validate_metadata_address(policy, "::ffff:169.254.169.254"), + mcp::auth::MetadataUrlDecision::address_link_local); + EXPECT_EQ(mcp::auth::validate_metadata_address(policy, "::ffff:10.0.0.1"), + mcp::auth::MetadataUrlDecision::address_private); +} + +TEST(AuthMetadataPolicyTest, AllowsRoutablePublicAddresses) { + const mcp::auth::MetadataFetchPolicy policy; + EXPECT_EQ(mcp::auth::validate_metadata_address(policy, "93.184.216.34"), + mcp::auth::MetadataUrlDecision::allowed); + EXPECT_EQ(mcp::auth::validate_metadata_address(policy, "2606:2800:220:1::1"), + mcp::auth::MetadataUrlDecision::allowed); +} + +TEST(AuthMetadataPolicyTest, DescribesEveryDecision) { + EXPECT_EQ(mcp::auth::describe(mcp::auth::MetadataUrlDecision::allowed), "allowed"); + for (const auto decision : {mcp::auth::MetadataUrlDecision::malformed_url, + mcp::auth::MetadataUrlDecision::scheme_not_allowed, + mcp::auth::MetadataUrlDecision::origin_denied, + mcp::auth::MetadataUrlDecision::origin_not_allowed, + mcp::auth::MetadataUrlDecision::address_link_local, + mcp::auth::MetadataUrlDecision::address_private, + mcp::auth::MetadataUrlDecision::address_loopback, + mcp::auth::MetadataUrlDecision::address_multicast, + mcp::auth::MetadataUrlDecision::address_reserved, + mcp::auth::MetadataUrlDecision::redirect_limit_exceeded, + mcp::auth::MetadataUrlDecision::response_too_large}) { + EXPECT_FALSE(mcp::auth::describe(decision).empty()); + EXPECT_NE(mcp::auth::describe(decision), "allowed"); + } +} + +TEST(AuthMetadataPolicyTest, PolicyErrorCarriesItsDecisionAndTarget) { + const mcp::auth::MetadataPolicyError error(mcp::auth::MetadataUrlDecision::address_link_local, + "169.254.169.254"); + EXPECT_EQ(error.decision(), mcp::auth::MetadataUrlDecision::address_link_local); + EXPECT_EQ(error.target(), "169.254.169.254"); + EXPECT_NE(std::string(error.what()).find("link-local"), std::string::npos); + EXPECT_NE(std::string(error.what()).find("169.254.169.254"), std::string::npos); +} + +TEST(AuthMetadataPolicyTest, ExtractsOriginsWithoutNormalizing) { + EXPECT_EQ(mcp::auth::metadata_url_origin("https://AS.test:8443/a/b?c=d"), "https://AS.test:8443"); + EXPECT_EQ(mcp::auth::metadata_url_origin("http://127.0.0.1:9000"), "http://127.0.0.1:9000"); + EXPECT_TRUE(mcp::auth::metadata_url_origin("not-a-url").empty()); +} diff --git a/test/auth/auth_client_identity_test.cpp b/test/auth/auth_client_identity_test.cpp new file mode 100644 index 0000000..8f723e8 --- /dev/null +++ b/test/auth/auth_client_identity_test.cpp @@ -0,0 +1,833 @@ +/** + * @file auth_client_identity_test.cpp + * @brief Tests for client identity selection: metadata documents, injected credentials, + * dynamic registration as the last resort, and the issuer binding that keeps one + * authorization server's credentials away from another. + */ + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace asio = boost::asio; +namespace beast = boost::beast; +namespace http = beast::http; +using json = nlohmann::json; + +namespace { + +/// Loopback HTTP server that answers a scripted handler and records every request it served. +class LoopbackServer final { + public: + using Handler = + std::function(const http::request&)>; + + explicit LoopbackServer(asio::io_context& io_ctx) + : acceptor_(io_ctx, {asio::ip::make_address("127.0.0.1"), 0}) {} + + [[nodiscard]] unsigned short port() const { return acceptor_.local_endpoint().port(); } + [[nodiscard]] std::string base_url() const { return "http://127.0.0.1:" + std::to_string(port()); } + [[nodiscard]] std::string origin() const { return base_url(); } + + void set_handler(Handler handler) { handler_ = std::move(handler); } + + [[nodiscard]] const std::vector& targets() const { return targets_; } + [[nodiscard]] const std::vector& bodies() const { return bodies_; } + + /// @return The `Authorization` header of the first request served for a target. + [[nodiscard]] std::string authorization_for(const std::string& target) const { + for (std::size_t index = 0; index < targets_.size(); ++index) { + if (targets_[index] == target) { + return authorizations_[index]; + } + } + return {}; + } + + /// @return How many requests were served for a target. + [[nodiscard]] std::size_t count(const std::string& target) const { + return static_cast(std::count(targets_.begin(), targets_.end(), target)); + } + + /// @return The body of the first request served for a target, or an empty string. + [[nodiscard]] std::string body_for(const std::string& target) const { + for (std::size_t index = 0; index < targets_.size(); ++index) { + if (targets_[index] == target) { + return bodies_[index]; + } + } + return {}; + } + + /// Serve at most `request_budget` requests, then stop. close() aborts a pending accept so the + /// io_context always drains. + asio::awaitable serve(int request_budget) { + for (int index = 0; index < request_budget; ++index) { + boost::system::error_code accept_error; + auto socket = co_await acceptor_.async_accept( + asio::redirect_error(asio::use_awaitable, accept_error)); + if (accept_error) { + co_return; + } + + beast::tcp_stream stream(std::move(socket)); + beast::flat_buffer buffer; + http::request request; + boost::system::error_code read_error; + co_await http::async_read(stream, buffer, request, + asio::redirect_error(asio::use_awaitable, read_error)); + if (read_error) { + co_return; + } + + targets_.emplace_back(request.target()); + bodies_.push_back(request.body()); + authorizations_.emplace_back(request[http::field::authorization]); + + auto response = handler_(request); + response.version(request.version()); + response.prepare_payload(); + boost::system::error_code write_error; + co_await http::async_write(stream, response, + asio::redirect_error(asio::use_awaitable, write_error)); + + beast::error_code shutdown_error; + (void)stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + } + } + + void close() { + boost::system::error_code ignored; + (void)acceptor_.close(ignored); + } + + private: + asio::ip::tcp::acceptor acceptor_; + Handler handler_; + std::vector targets_; + std::vector bodies_; + std::vector authorizations_; +}; + +http::response json_response(const json& body, + http::status status = http::status::ok) { + http::response response{status, 11}; + response.set(http::field::content_type, "application/json"); + response.body() = body.dump(); + return response; +} + +http::response not_found() { + http::response response{http::status::not_found, 11}; + response.body() = "{}"; + return response; +} + +json token_document() { + return {{"access_token", "granted-access-token"}, {"token_type", "Bearer"}, {"expires_in", 3600}}; +} + +mcp::auth::MetadataFetchPolicy loopback_policy(const std::string& origin) { + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins.push_back(origin); + policy.allow_plain_http_loopback = true; + return policy; +} + +mcp::auth::AuthorizationCallback echoing_callback(std::string* captured_url) { + return [captured_url](const mcp::auth::AuthorizationRequest& request) + -> mcp::Task { + if (captured_url != nullptr) { + *captured_url = request.authorization_url; + } + mcp::auth::AuthorizationResponse response; + response.code = "test-authorization-code"; + response.state = request.state; + response.iss = request.issuer; + co_return response; + }; +} + +mcp::auth::ClientIdentityServerFacts facts_with(std::string issuer, bool cimd_supported, + std::optional registration_endpoint) { + mcp::auth::ClientIdentityServerFacts facts; + facts.issuer = std::move(issuer); + facts.client_id_metadata_document_supported = cimd_supported; + facts.registration_endpoint = std::move(registration_endpoint); + return facts; +} + +mcp::auth::OAuthClientInformation credentials(std::string client_id, std::string issuer) { + mcp::auth::OAuthClientInformation information; + information.client_id = std::move(client_id); + information.issuer = std::move(issuer); + return information; +} + +} // namespace + +TEST(AuthClientIdentitySelectionTest, InjectedCredentialsWinOverEveryOtherPath) { + mcp::auth::ClientIdentityConfig config; + config.pre_registered = credentials("injected-client", "https://issuer.example"); + config.client_metadata_url = "https://client.example/metadata.json"; + + const auto decision = mcp::auth::select_client_identity( + config, facts_with("https://issuer.example", true, "https://issuer.example/register"), + credentials("stored-client", "https://issuer.example")); + + EXPECT_EQ(decision, mcp::auth::ClientIdentityDecision::use_pre_registered); +} + +TEST(AuthClientIdentitySelectionTest, InjectedCredentialsNeverFallBackToRegistration) { + mcp::auth::ClientIdentityConfig config; + config.pre_registered = credentials("injected-client", "https://issuer.example"); + + // Every other path is available and none of them is taken. + const auto decision = mcp::auth::select_client_identity( + config, facts_with("https://issuer.example", true, "https://issuer.example/register"), + std::nullopt); + + EXPECT_EQ(decision, mcp::auth::ClientIdentityDecision::use_pre_registered); +} + +TEST(AuthClientIdentitySelectionTest, PrefersTheMetadataDocumentWhenTheServerAdvertisesIt) { + mcp::auth::ClientIdentityConfig config; + config.client_metadata_url = "https://client.example/metadata.json"; + + const auto decision = mcp::auth::select_client_identity( + config, facts_with("https://issuer.example", true, "https://issuer.example/register"), + std::nullopt); + + EXPECT_EQ(decision, mcp::auth::ClientIdentityDecision::use_client_id_metadata_document); +} + +TEST(AuthClientIdentitySelectionTest, IgnoresTheMetadataDocumentWhenTheServerDoesNotAdvertiseIt) { + mcp::auth::ClientIdentityConfig config; + config.client_metadata_url = "https://client.example/metadata.json"; + + const auto decision = mcp::auth::select_client_identity( + config, facts_with("https://issuer.example", false, "https://issuer.example/register"), + std::nullopt); + + EXPECT_EQ(decision, mcp::auth::ClientIdentityDecision::register_dynamically); +} + +TEST(AuthClientIdentitySelectionTest, ReusesCredentialsRecordedAgainstTheSameIssuer) { + const auto decision = mcp::auth::select_client_identity( + {}, facts_with("https://issuer.example", false, "https://issuer.example/register"), + credentials("stored-client", "https://issuer.example")); + + EXPECT_EQ(decision, mcp::auth::ClientIdentityDecision::reuse_stored_registration); +} + +TEST(AuthClientIdentitySelectionTest, RegistersAfreshWhenTheStoredIssuerIsADifferentServer) { + const auto decision = mcp::auth::select_client_identity( + {}, facts_with("https://second.example", false, "https://second.example/register"), + credentials("stored-client", "https://first.example")); + + EXPECT_EQ(decision, mcp::auth::ClientIdentityDecision::register_dynamically); +} + +TEST(AuthClientIdentitySelectionTest, ReportsUnavailableWhenNoPathRemains) { + const auto decision = mcp::auth::select_client_identity( + {}, facts_with("https://issuer.example", true, std::nullopt), std::nullopt); + + EXPECT_EQ(decision, mcp::auth::ClientIdentityDecision::unavailable); +} + +TEST(AuthClientRegistrationRequestTest, CarriesApplicationTypeAndTheRefreshTokenGrant) { + mcp::auth::OAuthClientMetadata metadata; + metadata.redirect_uris = {"http://127.0.0.1:9999/callback"}; + metadata.client_name = "conformance-client"; + + const auto body = mcp::auth::build_registration_request( + metadata, facts_with("https://issuer.example", false, "https://issuer.example/register")); + + EXPECT_EQ(body.at("application_type"), "native"); + EXPECT_EQ(body.at("grant_types"), json::array({"authorization_code", "refresh_token"})); + EXPECT_EQ(body.at("response_types"), json::array({"code"})); + EXPECT_EQ(body.at("redirect_uris"), json::array({"http://127.0.0.1:9999/callback"})); + EXPECT_EQ(body.at("client_name"), "conformance-client"); + EXPECT_FALSE(body.contains("scope")); +} + +TEST(AuthClientRegistrationRequestTest, RequestsOfflineAccessOnlyWhenTheServerAdvertisesIt) { + mcp::auth::OAuthClientMetadata metadata; + metadata.redirect_uris = {"http://127.0.0.1:9999/callback"}; + metadata.scope = "mcp:read"; + + auto silent = facts_with("https://issuer.example", false, "https://issuer.example/register"); + silent.scopes_supported = {"mcp:read", "mcp:write"}; + EXPECT_EQ(mcp::auth::build_registration_request(metadata, silent).at("scope"), "mcp:read"); + + auto advertising = silent; + advertising.scopes_supported.emplace_back("offline_access"); + EXPECT_EQ(mcp::auth::build_registration_request(metadata, advertising).at("scope"), + "mcp:read offline_access"); +} + +TEST(AuthClientRegistrationRequestTest, DoesNotRepeatOfflineAccessAlreadyRequested) { + mcp::auth::OAuthClientMetadata metadata; + metadata.scope = "mcp:read offline_access"; + + auto facts = facts_with("https://issuer.example", false, "https://issuer.example/register"); + facts.scopes_supported = {"offline_access"}; + + EXPECT_EQ(mcp::auth::build_registration_request(metadata, facts).at("scope"), + "mcp:read offline_access"); +} + +TEST(AuthClientInformationTest, ReportsSecretExpiryOnlyForANonZeroExpiryInThePast) { + auto information = credentials("client", "https://issuer.example"); + EXPECT_FALSE(information.secret_expired(1000)); + + information.client_secret_expires_at = 0; // RFC 7591: zero means the secret never expires. + EXPECT_FALSE(information.secret_expired(1000)); + + information.client_secret_expires_at = 2000; + EXPECT_FALSE(information.secret_expired(1000)); + + information.client_secret_expires_at = 999; + EXPECT_TRUE(information.secret_expired(1000)); +} + +TEST(AuthClientInformationTest, SerializesTheIssuerBindingAlongsideTheCredentials) { + auto information = credentials("client", "https://issuer.example"); + information.client_secret = "secret"; + information.client_secret_expires_at = 1234; + + const nlohmann::json serialized = information; + EXPECT_EQ(serialized.at("client_id"), "client"); + EXPECT_EQ(serialized.at("client_secret"), "secret"); + EXPECT_EQ(serialized.at("client_secret_expires_at"), 1234); + EXPECT_EQ(serialized.at("issuer"), "https://issuer.example"); + + // The wire format has no issuer, so a round trip recovers everything the server sent and the + // caller re-applies the binding. + const auto parsed = serialized.get(); + EXPECT_EQ(parsed.client_id, "client"); + EXPECT_EQ(parsed.client_secret, "secret"); + EXPECT_EQ(parsed.client_secret_expires_at, 1234); + EXPECT_TRUE(parsed.issuer.empty()); +} + +TEST(AuthClientIdentitySelectionTest, DescribesEverySourceAndDecision) { + EXPECT_FALSE(mcp::auth::describe(mcp::auth::ClientIdentitySource::pre_registered).empty()); + EXPECT_FALSE( + mcp::auth::describe(mcp::auth::ClientIdentitySource::client_id_metadata_document).empty()); + EXPECT_FALSE(mcp::auth::describe(mcp::auth::ClientIdentitySource::dynamic_registration).empty()); + EXPECT_FALSE(mcp::auth::describe(mcp::auth::ClientIdentityDecision::use_pre_registered).empty()); + EXPECT_FALSE(mcp::auth::describe(mcp::auth::ClientIdentityDecision::use_client_id_metadata_document) + .empty()); + EXPECT_FALSE( + mcp::auth::describe(mcp::auth::ClientIdentityDecision::reuse_stored_registration).empty()); + EXPECT_FALSE(mcp::auth::describe(mcp::auth::ClientIdentityDecision::register_dynamically).empty()); + EXPECT_FALSE(mcp::auth::describe(mcp::auth::ClientIdentityDecision::unavailable).empty()); +} + +TEST(AuthMetadataOriginAllowanceTest, WidensTheAllowListWithoutOverridingTheDenyList) { + mcp::auth::MetadataFetchPolicy policy; + policy.allow_plain_http_loopback = true; + policy.denied_origins.emplace_back("http://127.0.0.1:9"); + + // With no hook, an unlisted origin is refused; the hook admits exactly what it says yes to. + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://127.0.0.1:1234/prm"), + mcp::auth::MetadataUrlDecision::origin_not_allowed); + + policy.origin_allowance = [](const std::string& origin) { + return origin == "http://127.0.0.1:1234"; + }; + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://127.0.0.1:1234/prm"), + mcp::auth::MetadataUrlDecision::allowed); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://127.0.0.1:5678/prm"), + mcp::auth::MetadataUrlDecision::origin_not_allowed); + + // A denied origin stays denied even when the hook would admit it. + policy.origin_allowance = [](const std::string&) { return true; }; + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://127.0.0.1:9/prm"), + mcp::auth::MetadataUrlDecision::origin_denied); + + // Widening the origin list does not relax any other control: the scheme rule still refuses + // plain HTTP to a non-loopback host, and the address rules still refuse the metadata service. + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "http://169.254.169.254/prm"), + mcp::auth::MetadataUrlDecision::scheme_not_allowed); + EXPECT_EQ(mcp::auth::validate_metadata_url(policy, "https://169.254.169.254/prm"), + mcp::auth::MetadataUrlDecision::address_link_local); +} + +TEST(AuthClientCredentialStoreTest, KeepsEachIssuersCredentialsSeparate) { + mcp::auth::InMemoryClientCredentialStore store; + store.store("https://first.example", credentials("client-1", "https://first.example")); + store.store("https://second.example", credentials("client-2", "https://second.example")); + + ASSERT_TRUE(store.load("https://first.example").has_value()); + EXPECT_EQ(store.load("https://first.example")->client_id, "client-1"); + EXPECT_EQ(store.load("https://second.example")->client_id, "client-2"); + EXPECT_FALSE(store.load("https://third.example").has_value()); + + store.remove("https://first.example"); + EXPECT_FALSE(store.load("https://first.example").has_value()); + EXPECT_TRUE(store.load("https://second.example").has_value()); +} + +namespace { + +/// Authorization-server metadata for a fixture server rooted at `base` with issuer `base + suffix`. +json server_metadata(const std::string& base, const std::string& suffix, bool cimd_supported, + bool registration_supported) { + json metadata = {{"issuer", base + suffix}, + {"authorization_endpoint", base + suffix + "/authorize"}, + {"token_endpoint", base + suffix + "/token"}, + {"response_types_supported", json::array({"code"})}, + {"code_challenge_methods_supported", json::array({"S256"})}}; + if (cimd_supported) { + metadata["client_id_metadata_document_supported"] = true; + } + if (registration_supported) { + metadata["registration_endpoint"] = base + suffix + "/register"; + } + return metadata; +} + +struct IdentityFixture { + std::shared_ptr tokens = + std::make_shared(); + std::shared_ptr credentials = + std::make_shared(); + mcp::auth::OAuthAuthorizationConfig config; + std::string authorization_url; + std::optional identity; + std::exception_ptr failure; +}; + +/// Run one challenge against a single-origin fixture and capture what the manager chose. +void run_single_challenge(LoopbackServer& server, IdentityFixture& fixture, asio::io_context& io_ctx, + const std::string& challenge, int request_budget) { + asio::co_spawn(io_ctx, server.serve(request_budget), asio::detached); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.tokens, + fixture.config, + echoing_callback(&fixture.authorization_url)); + try { + (void)co_await manager.try_handle_challenge(challenge); + fixture.identity = manager.last_client_identity(); + } catch (...) { + fixture.failure = std::current_exception(); + } + server.close(); + }, + asio::detached); + io_ctx.run(); +} + +} // namespace + +TEST(AuthClientIdentityLoopbackTest, UsesTheMetadataDocumentUrlAsTheClientIdAndSkipsRegistration) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(server_metadata(base, "", true, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return not_found(); + }); + + IdentityFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.client_identity.client_metadata_url = "https://client.example/metadata.json"; + fixture.config.credential_store = fixture.credentials; + fixture.config.policy = loopback_policy(server.origin()); + + run_single_challenge(server, fixture, io_ctx, R"(Bearer resource_metadata=")" + base + R"(/prm")", + 3); + + ASSERT_EQ(fixture.failure, nullptr); + ASSERT_TRUE(fixture.identity.has_value()); + EXPECT_EQ(fixture.identity->source, mcp::auth::ClientIdentitySource::client_id_metadata_document); + EXPECT_EQ(fixture.identity->client_id, "https://client.example/metadata.json"); + EXPECT_FALSE(fixture.identity->client_secret.has_value()); + + // The registration endpoint was advertised and deliberately not used. + EXPECT_EQ(server.count("/register"), 0U); + EXPECT_NE(fixture.authorization_url.find( + mcp::auth::detail::url_encode("https://client.example/metadata.json")), + std::string::npos); + EXPECT_NE(server.body_for("/token").find( + "client_id=" + mcp::auth::detail::url_encode("https://client.example/metadata.json")), + std::string::npos); + + // Nothing is persisted: the document URL is the identity, so there is no credential to bind. + EXPECT_FALSE(fixture.credentials->load(base).has_value()); +} + +TEST(AuthClientIdentityLoopbackTest, PreRegisteredCredentialsAreNeverExchangedForARegistration) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + // Registration is available and the server also advertises metadata documents. + return json_response(server_metadata(base, "", true, true)); + } + if (target == "/token") { + return json_response(token_document()); + } + return not_found(); + }); + + IdentityFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + // Bound to the issuer these credentials were registered with, which is the precondition for a + // confidential client: the SDK refuses to present a secret that names no issuer. + fixture.config.client_identity.pre_registered = credentials("application-chosen-client", base); + fixture.config.client_identity.pre_registered->client_secret = "application-chosen-secret"; + fixture.config.client_identity.client_metadata_url = "https://client.example/metadata.json"; + fixture.config.credential_store = fixture.credentials; + fixture.config.policy = loopback_policy(server.origin()); + + run_single_challenge(server, fixture, io_ctx, R"(Bearer resource_metadata=")" + base + R"(/prm")", + 3); + + ASSERT_EQ(fixture.failure, nullptr); + ASSERT_TRUE(fixture.identity.has_value()); + EXPECT_EQ(fixture.identity->source, mcp::auth::ClientIdentitySource::pre_registered); + EXPECT_EQ(fixture.identity->client_id, "application-chosen-client"); + + EXPECT_EQ(server.count("/register"), 0U); + EXPECT_NE(server.body_for("/token").find("client_id=application-chosen-client"), std::string::npos); + EXPECT_NE(server.body_for("/token").find("client_secret=application-chosen-secret"), + std::string::npos); + // Injected credentials belong to the application; the SDK does not persist them. + EXPECT_FALSE(fixture.credentials->load(base).has_value()); +} + +TEST(AuthClientIdentityLoopbackTest, FallsBackToDynamicRegistrationAndSendsApplicationType) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + auto metadata = server_metadata(base, "", false, true); + metadata["scopes_supported"] = json::array({"mcp:read", "offline_access"}); + return json_response(metadata); + } + if (target == "/register") { + return json_response({{"client_id", "registered-client"}, + {"client_secret", "registered-secret"}, + {"client_secret_expires_at", 0}}, + http::status::created); + } + if (target == "/token") { + return json_response(token_document()); + } + return not_found(); + }); + + IdentityFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.client_identity.metadata.client_name = "mcp-cpp-sdk"; + fixture.config.credential_store = fixture.credentials; + fixture.config.policy = loopback_policy(server.origin()); + + run_single_challenge(server, fixture, io_ctx, R"(Bearer resource_metadata=")" + base + R"(/prm")", + 4); + + ASSERT_EQ(fixture.failure, nullptr); + ASSERT_TRUE(fixture.identity.has_value()); + EXPECT_EQ(fixture.identity->source, mcp::auth::ClientIdentitySource::dynamic_registration); + EXPECT_EQ(fixture.identity->client_id, "registered-client"); + + ASSERT_EQ(server.count("/register"), 1U); + const auto registration = json::parse(server.body_for("/register")); + EXPECT_EQ(registration.at("application_type"), "native"); + EXPECT_EQ(registration.at("grant_types"), json::array({"authorization_code", "refresh_token"})); + // The redirect URI configured for authorization is the one registered. + EXPECT_EQ(registration.at("redirect_uris"), json::array({"http://127.0.0.1:9999/callback"})); + EXPECT_EQ(registration.at("scope"), "offline_access"); + + EXPECT_NE(server.body_for("/token").find("client_id=registered-client"), std::string::npos); + + const auto stored = fixture.credentials->load(base); + ASSERT_TRUE(stored.has_value()); + EXPECT_EQ(stored->client_id, "registered-client"); + EXPECT_EQ(stored->issuer, base); +} + +TEST(AuthClientIdentityLoopbackTest, KeepsCredentialsBoundToTheIssuerThatGrantedThem) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm-one") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base + "/as1"})}}); + } + if (target == "/prm-two") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base + "/as2"})}}); + } + if (target == "/.well-known/oauth-authorization-server/as1") { + return json_response(server_metadata(base, "/as1", false, true)); + } + if (target == "/.well-known/oauth-authorization-server/as2") { + return json_response(server_metadata(base, "/as2", false, true)); + } + if (target == "/as1/register") { + return json_response({{"client_id", "client-for-as-1"}, {"client_secret", "secret-one"}}, + http::status::created); + } + if (target == "/as2/register") { + return json_response({{"client_id", "client-for-as-2"}, {"client_secret", "secret-two"}}, + http::status::created); + } + if (target == "/as1/token" || target == "/as2/token") { + return json_response(token_document()); + } + return not_found(); + }); + asio::co_spawn(io_ctx, server.serve(8), asio::detached); + + IdentityFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.credential_store = fixture.credentials; + fixture.config.policy = loopback_policy(server.origin()); + + std::optional first_identity; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.tokens, + fixture.config, echoing_callback(nullptr)); + try { + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm-one")"); + first_identity = manager.last_client_identity(); + // The protected resource now names a different authorization server. + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm-two")"); + fixture.identity = manager.last_client_identity(); + } catch (...) { + fixture.failure = std::current_exception(); + } + server.close(); + }, + asio::detached); + io_ctx.run(); + + ASSERT_EQ(fixture.failure, nullptr); + ASSERT_TRUE(first_identity.has_value()); + EXPECT_EQ(first_identity->client_id, "client-for-as-1"); + EXPECT_EQ(first_identity->issuer, base + "/as1"); + + // The authorization server changed, so a fresh registration is performed rather than the first + // server's credentials being reused. + ASSERT_TRUE(fixture.identity.has_value()); + EXPECT_EQ(fixture.identity->client_id, "client-for-as-2"); + EXPECT_EQ(fixture.identity->issuer, base + "/as2"); + EXPECT_EQ(server.count("/as1/register"), 1U); + EXPECT_EQ(server.count("/as2/register"), 1U); + + // Nothing the second server received mentions the first server's client. + for (std::size_t index = 0; index < server.targets().size(); ++index) { + if (server.targets()[index].rfind("/as2/", 0) == 0) { + EXPECT_EQ(server.bodies()[index].find("client-for-as-1"), std::string::npos) + << "AS-2 was sent AS-1's client identifier on " << server.targets()[index]; + } + } + + const auto first = fixture.credentials->load(base + "/as1"); + const auto second = fixture.credentials->load(base + "/as2"); + ASSERT_TRUE(first.has_value()); + ASSERT_TRUE(second.has_value()); + EXPECT_EQ(first->client_id, "client-for-as-1"); + EXPECT_EQ(second->client_id, "client-for-as-2"); +} + +TEST(AuthTokenEndpointAuthLoopbackTest, PutsTheSecretInHttpBasicWhenTheServerAsksForIt) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + auto metadata = server_metadata(base, "", false, false); + metadata["token_endpoint_auth_methods_supported"] = json::array({"client_secret_basic"}); + return json_response(metadata); + } + if (target == "/token") { + return json_response(token_document()); + } + return not_found(); + }); + + IdentityFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.client_id = "confidential-client"; + fixture.config.client_secret = "confidential-secret"; + // A confidential client must name the issuer its secret belongs to; without it the SDK refuses + // to present the secret at all, which is what makes the binding guard load-bearing. + fixture.config.client_issuer = base; + fixture.config.policy = loopback_policy(server.origin()); + + run_single_challenge(server, fixture, io_ctx, R"(Bearer resource_metadata=")" + base + R"(/prm")", + 3); + + ASSERT_EQ(fixture.failure, nullptr); + // "confidential-client:confidential-secret" base64-encoded. + EXPECT_EQ(server.authorization_for("/token"), + "Basic Y29uZmlkZW50aWFsLWNsaWVudDpjb25maWRlbnRpYWwtc2VjcmV0"); + EXPECT_EQ(server.body_for("/token").find("client_secret"), std::string::npos); +} + +TEST(AuthTokenEndpointAuthLoopbackTest, SendsNoSecretAtAllWhenTheServerAdvertisesOnlyNone) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + auto metadata = server_metadata(base, "", false, true); + metadata["token_endpoint_auth_methods_supported"] = json::array({"none"}); + return json_response(metadata); + } + if (target == "/register") { + return json_response( + {{"client_id", "registered-client"}, {"client_secret", "registered-secret"}}, + http::status::created); + } + if (target == "/token") { + return json_response(token_document()); + } + return not_found(); + }); + + IdentityFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.credential_store = fixture.credentials; + fixture.config.policy = loopback_policy(server.origin()); + + run_single_challenge(server, fixture, io_ctx, R"(Bearer resource_metadata=")" + base + R"(/prm")", + 4); + + ASSERT_EQ(fixture.failure, nullptr); + ASSERT_TRUE(fixture.identity.has_value()); + // Registration handed the client a secret the server then refuses to accept anywhere. + EXPECT_EQ(fixture.identity->client_secret, "registered-secret"); + EXPECT_TRUE(server.authorization_for("/token").empty()); + EXPECT_EQ(server.body_for("/token").find("client_secret"), std::string::npos); + EXPECT_EQ(server.body_for("/token").find("registered-secret"), std::string::npos); +} + +TEST(AuthClientIdentityLoopbackTest, ReusesTheRegistrationRecordedForTheSameIssuer) { + asio::io_context io_ctx; + LoopbackServer server(io_ctx); + const auto base = server.base_url(); + + server.set_handler([&base](const http::request& request) { + const std::string target(request.target()); + if (target == "/prm-one" || target == "/prm-two") { + return json_response( + {{"resource", base + "/mcp"}, {"authorization_servers", json::array({base})}}); + } + if (target == "/.well-known/oauth-authorization-server") { + return json_response(server_metadata(base, "", false, true)); + } + if (target == "/register") { + return json_response({{"client_id", "registered-client"}}, http::status::created); + } + if (target == "/token") { + return json_response(token_document()); + } + return not_found(); + }); + asio::co_spawn(io_ctx, server.serve(7), asio::detached); + + IdentityFixture fixture; + fixture.config.server_url = base + "/mcp"; + fixture.config.redirect_uri = "http://127.0.0.1:9999/callback"; + fixture.config.credential_store = fixture.credentials; + fixture.config.policy = loopback_policy(server.origin()); + + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::auth::OAuthAuthorizationManager manager(io_ctx.get_executor(), fixture.tokens, + fixture.config, echoing_callback(nullptr)); + try { + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm-one")"); + (void)co_await manager.try_handle_challenge(R"(Bearer resource_metadata=")" + base + + R"(/prm-two")"); + fixture.identity = manager.last_client_identity(); + } catch (...) { + fixture.failure = std::current_exception(); + } + server.close(); + }, + asio::detached); + io_ctx.run(); + + ASSERT_EQ(fixture.failure, nullptr); + ASSERT_TRUE(fixture.identity.has_value()); + EXPECT_EQ(fixture.identity->source, mcp::auth::ClientIdentitySource::dynamic_registration); + EXPECT_EQ(fixture.identity->client_id, "registered-client"); + EXPECT_EQ(server.count("/register"), 1U); +} diff --git a/test/auth/auth_integration_test.cpp b/test/auth/auth_integration_test.cpp index 4507dea..9e75a0f 100644 --- a/test/auth/auth_integration_test.cpp +++ b/test/auth/auth_integration_test.cpp @@ -11,11 +11,16 @@ #include #include #include +#include +#include #include #include +#include #include #include #include +#include +#include #include #include #include @@ -26,6 +31,56 @@ namespace beast = boost::beast; namespace http = beast::http; using json = nlohmann::json; +namespace { + +class RotatingAuthenticator final : public mcp::auth::Authenticator { + public: + [[nodiscard]] std::string get_access_token() const override { return access_token_; } + + mcp::Task try_refresh_token() override { + ++refresh_count_; + access_token_ = "refreshed-token"; + co_return true; + } + + [[nodiscard]] int refresh_count() const { return refresh_count_; } + + private: + std::string access_token_{"initial-token"}; + int refresh_count_{0}; +}; + +class CloseTransparentTransport final : public mcp::ITransport { + public: + mcp::Task read_message() override { + if (incoming_.empty()) { + throw std::runtime_error("no scripted response"); + } + auto message = std::move(incoming_.front()); + incoming_.pop(); + co_return message; + } + + mcp::Task write_message(std::string_view message) override { + written_.emplace_back(message); + co_return; + } + + void close() override { ++close_count_; } + + void enqueue_message(std::string message) { incoming_.push(std::move(message)); } + + [[nodiscard]] const std::vector& written() const { return written_; } + [[nodiscard]] int close_count() const { return close_count_; } + + private: + std::queue incoming_; + std::vector written_; + int close_count_{0}; +}; + +} // namespace + TEST(AuthBearerExtractionTest, ValidBearerToken) { auto token = mcp::auth::extract_bearer_token("Bearer my_access_token"); EXPECT_EQ(token, "my_access_token"); @@ -88,6 +143,7 @@ TEST(AuthMiddlewareTest, AcceptsValidToken) { }); transport_ptr->enqueue_message(make_initialize_request("1").dump()); + transport_ptr->enqueue_message(make_initialized_notification().dump()); json call_req{{"jsonrpc", "2.0"}, {"id", "2"}, @@ -139,6 +195,7 @@ TEST(AuthMiddlewareTest, RejectsInvalidToken) { }); transport_ptr->enqueue_message(make_initialize_request("1").dump()); + transport_ptr->enqueue_message(make_initialized_notification().dump()); json call_req{{"jsonrpc", "2.0"}, {"id", "2"}, @@ -192,6 +249,7 @@ TEST(AuthMiddlewareTest, RejectsMissingToken) { }); transport_ptr->enqueue_message(make_initialize_request("1").dump()); + transport_ptr->enqueue_message(make_initialized_notification().dump()); json call_req{{"jsonrpc", "2.0"}, {"id", "2"}, @@ -276,6 +334,299 @@ TEST(AuthClientTransportTest, ReadWritePassThrough) { transport.close(); } +TEST(AuthClientTransportTest, LegacyRetryReplaysTheRequestMatchingTheErrorId) { + asio::io_context io; + auto inner = std::make_shared(io.get_executor()); + auto authenticator = std::make_shared(); + mcp::auth::OAuthClientTransport transport(inner, authenticator); + + const auto request_a = + json{{"jsonrpc", "2.0"}, {"id", "A"}, {"method", "tools/call"}, {"params", json::object()}} + .dump(); + const auto request_b = + json{{"jsonrpc", "2.0"}, {"id", "B"}, {"method", "tools/call"}, {"params", json::object()}} + .dump(); + inner->enqueue_message(make_error_response("A", mcp::g_UNAUTHORIZED, "Unauthorized").dump()); + inner->enqueue_message(make_result_response("A", json{{"ok", true}}).dump()); + + std::string response; + asio::co_spawn( + io, + [&]() -> mcp::Task { + co_await transport.write_message(request_a); + co_await transport.write_message(request_b); + response = co_await transport.read_message(); + }, + asio::detached); + io.run(); + + ASSERT_EQ(inner->written().size(), 3); + const auto first_request = json::parse(inner->written()[0]); + const auto second_request = json::parse(inner->written()[1]); + const auto retried_request = json::parse(inner->written()[2]); + EXPECT_EQ(first_request["id"], "A"); + EXPECT_EQ(second_request["id"], "B"); + EXPECT_EQ(retried_request["id"], "A"); + EXPECT_EQ(first_request["params"]["_meta"]["auth_token"], "initial-token"); + EXPECT_EQ(retried_request["params"]["_meta"]["auth_token"], "refreshed-token"); + EXPECT_EQ(authenticator->refresh_count(), 1); + EXPECT_EQ(json::parse(response)["id"], "A"); +} + +TEST(AuthClientTransportTest, NoResponseRequestsAreEvictedAtConfiguredBound) { + asio::io_context io; + auto inner = std::make_shared(io.get_executor()); + auto authenticator = std::make_shared(); + mcp::auth::OAuthClientTransportOptions options; + options.max_pending_requests = 2; + options.pending_request_ttl = std::chrono::hours(1); + mcp::auth::OAuthClientTransport transport(inner, authenticator, options); + + const auto make_request = [](std::string_view id) { + return json{{"jsonrpc", "2.0"}, {"id", id}, {"method", "tools/list"}}.dump(); + }; + inner->enqueue_message(make_error_response("A", mcp::g_UNAUTHORIZED, "Unauthorized").dump()); + inner->enqueue_message(make_error_response("C", mcp::g_UNAUTHORIZED, "Unauthorized").dump()); + inner->enqueue_message(make_result_response("C", json{{"ok", true}}).dump()); + + std::string evicted_response; + std::string retained_response; + asio::co_spawn( + io, + [&]() -> mcp::Task { + co_await transport.write_message(make_request("A")); + co_await transport.write_message(make_request("B")); + co_await transport.write_message(make_request("C")); + evicted_response = co_await transport.read_message(); + retained_response = co_await transport.read_message(); + }, + asio::detached); + io.run(); + + ASSERT_EQ(inner->written().size(), 4); + EXPECT_EQ(json::parse(evicted_response).at("id"), "A"); + EXPECT_EQ(json::parse(retained_response).at("id"), "C"); + EXPECT_EQ(json::parse(inner->written().back()).at("id"), "C"); + EXPECT_EQ(authenticator->refresh_count(), 1); +} + +TEST(AuthClientTransportTest, PendingReplayCorrelationExpires) { + asio::io_context io; + auto inner = std::make_shared(io.get_executor()); + auto authenticator = std::make_shared(); + mcp::auth::OAuthClientTransportOptions options; + options.pending_request_ttl = std::chrono::milliseconds(5); + mcp::auth::OAuthClientTransport transport(inner, authenticator, options); + + const auto request = json{{"jsonrpc", "2.0"}, {"id", "expired"}, {"method", "tools/list"}}.dump(); + inner->enqueue_message(make_error_response("expired", mcp::g_UNAUTHORIZED, "Unauthorized").dump()); + + std::string response; + asio::co_spawn( + io, + [&]() -> mcp::Task { + co_await transport.write_message(request); + asio::steady_timer expiry_wait(io); + expiry_wait.expires_after(std::chrono::milliseconds(25)); + co_await expiry_wait.async_wait(asio::use_awaitable); + response = co_await transport.read_message(); + }, + asio::detached); + io.run(); + + ASSERT_EQ(inner->written().size(), 1); + EXPECT_EQ(json::parse(response).at("id"), "expired"); + EXPECT_EQ(authenticator->refresh_count(), 0); +} + +TEST(AuthClientTransportTest, CloseClearsPendingReplayCorrelationAndIsIdempotent) { + asio::io_context io; + auto inner = std::make_shared(); + auto authenticator = std::make_shared(); + mcp::auth::OAuthClientTransport transport(inner, authenticator); + + const auto request = json{{"jsonrpc", "2.0"}, {"id", "closed"}, {"method", "tools/list"}}.dump(); + inner->enqueue_message(make_error_response("closed", mcp::g_UNAUTHORIZED, "Unauthorized").dump()); + + std::string response; + asio::co_spawn( + io, + [&]() -> mcp::Task { + co_await transport.write_message(request); + transport.close(); + transport.close(); + response = co_await transport.read_message(); + }, + asio::detached); + io.run(); + + ASSERT_EQ(inner->written().size(), 1); + EXPECT_EQ(inner->close_count(), 1); + EXPECT_EQ(json::parse(response).at("id"), "closed"); + EXPECT_EQ(authenticator->refresh_count(), 0); +} + +TEST(AuthClientTransportTest, HttpTransportUsesAuthorizationHeader) { + constexpr unsigned short port = 18109; + asio::io_context io; + asio::ip::tcp::acceptor acceptor(io, {asio::ip::make_address("127.0.0.1"), port}); + + std::string authorization_header; + asio::co_spawn( + io, + [&]() -> asio::awaitable { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + beast::tcp_stream stream(std::move(socket)); + beast::flat_buffer buffer; + http::request request; + co_await http::async_read(stream, buffer, request, asio::use_awaitable); + authorization_header = std::string(request[http::field::authorization]); + + http::response response{http::status::accepted, request.version()}; + response.content_length(0); + co_await http::async_write(stream, response, asio::use_awaitable); + }, + asio::detached); + + auto inner = std::make_shared( + io.get_executor(), "http://127.0.0.1:" + std::to_string(port) + "/mcp"); + auto store = std::make_shared(); + auto oauth_client = std::make_shared(io.get_executor()); + + mcp::auth::OAuthConfig config; + config.client_id = "test"; + config.token_endpoint = "http://localhost/token"; + config.redirect_uri = "http://localhost/callback"; + + mcp::auth::TokenResponse token; + token.access_token = "header-token"; + store->store("http://server1", token); + + auto authenticator = + std::make_shared(store, oauth_client, config, "http://server1"); + mcp::auth::OAuthClientTransport transport(inner, authenticator); + + asio::co_spawn( + io, + [&]() -> mcp::Task { + co_await transport.write_message(R"({"jsonrpc":"2.0","id":1,"method":"ping"})"); + transport.close(); + }, + asio::detached); + + io.run(); + + EXPECT_EQ(authorization_header, "Bearer header-token"); +} + +TEST(AuthClientTransportTest, FailedHttpRetryDoesNotLeaveRequestEligibleForReplay) { + asio::io_context io; + asio::ip::tcp::acceptor initial_acceptor(io); + initial_acceptor.open(asio::ip::tcp::v4()); + initial_acceptor.set_option(asio::socket_base::reuse_address(true)); + initial_acceptor.bind({asio::ip::make_address("127.0.0.1"), 0}); + initial_acceptor.listen(); + const auto port = initial_acceptor.local_endpoint().port(); + + asio::co_spawn( + io, + [&]() -> asio::awaitable { + auto socket = co_await initial_acceptor.async_accept(asio::use_awaitable); + initial_acceptor.close(); + + beast::tcp_stream stream(std::move(socket)); + beast::flat_buffer buffer; + http::request request; + co_await http::async_read(stream, buffer, request, asio::use_awaitable); + + http::response response{http::status::unauthorized, request.version()}; + response.set(http::field::www_authenticate, "Bearer"); + response.keep_alive(false); + response.content_length(0); + co_await http::async_write(stream, response, asio::use_awaitable); + }, + asio::detached); + + auto inner = std::make_shared( + io.get_executor(), "http://127.0.0.1:" + std::to_string(port) + "/mcp"); + auto authenticator = std::make_shared(); + mcp::auth::OAuthClientTransport transport(inner, authenticator); + + const auto original_request = + json{{"jsonrpc", "2.0"}, {"id", "A"}, {"method", "tools/list"}}.dump(); + const auto probe_notification = json{{"jsonrpc", "2.0"}, {"method", "notifications/probe"}}.dump(); + + bool retry_failed = false; + std::string response_wire; + std::vector replay_server_requests; + std::exception_ptr controller_error; + asio::co_spawn( + io, + [&]() -> mcp::Task { + try { + co_await transport.write_message(original_request); + } catch (const std::exception&) { + retry_failed = true; + } + + auto replay_acceptor = std::make_shared(io); + replay_acceptor->open(asio::ip::tcp::v4()); + replay_acceptor->set_option(asio::socket_base::reuse_address(true)); + replay_acceptor->bind({asio::ip::make_address("127.0.0.1"), port}); + replay_acceptor->listen(); + + asio::co_spawn( + io, + [replay_acceptor, &replay_server_requests]() -> asio::awaitable { + for (int request_index = 0; request_index < 2; ++request_index) { + boost::system::error_code accept_error; + auto socket = co_await replay_acceptor->async_accept( + asio::redirect_error(asio::use_awaitable, accept_error)); + if (accept_error == asio::error::operation_aborted) { + co_return; + } + if (accept_error) { + throw boost::system::system_error(accept_error); + } + + beast::tcp_stream stream(std::move(socket)); + beast::flat_buffer buffer; + http::request request; + co_await http::async_read(stream, buffer, request, asio::use_awaitable); + replay_server_requests.push_back(json::parse(request.body())); + + const auto response_body = + request_index == 0 + ? make_error_response("A", mcp::g_UNAUTHORIZED, "Unauthorized").dump() + : make_result_response("A", json{{"unexpectedReplay", true}}).dump(); + http::response response{http::status::ok, request.version()}; + response.set(http::field::content_type, "application/json"); + response.keep_alive(false); + response.body() = response_body; + response.prepare_payload(); + co_await http::async_write(stream, response, asio::use_awaitable); + } + }, + asio::detached); + + co_await transport.write_message(probe_notification); + response_wire = co_await transport.read_message(); + replay_acceptor->close(); + transport.close(); + }, + [&controller_error](std::exception_ptr error) { controller_error = std::move(error); }); + + io.run(); + + EXPECT_EQ(controller_error, nullptr); + EXPECT_TRUE(retry_failed); + EXPECT_EQ(authenticator->refresh_count(), 1); + ASSERT_EQ(replay_server_requests.size(), 1); + EXPECT_EQ(replay_server_requests.front().at("method"), "notifications/probe"); + ASSERT_FALSE(response_wire.empty()); + EXPECT_EQ(json::parse(response_wire).at("error").at("code"), mcp::g_UNAUTHORIZED); +} + TEST(AuthClientTransportTest, RefreshTokenReturnsTrue) { constexpr unsigned short port = 18107; asio::io_context io; @@ -311,6 +662,10 @@ TEST(AuthClientTransportTest, RefreshTokenReturnsTrue) { auto inner = std::make_shared(io.get_executor()); auto store = std::make_shared(); auto oauth_client = std::make_shared(io.get_executor()); + mcp::auth::MetadataFetchPolicy oauth_policy; + oauth_policy.allowed_origins.push_back("http://127.0.0.1:" + std::to_string(port)); + oauth_policy.allow_plain_http_loopback = true; + oauth_client->set_metadata_policy(oauth_policy); mcp::auth::OAuthConfig config; config.client_id = "test"; @@ -428,6 +783,10 @@ TEST(AuthClientTransportTest, RefreshPreservesOldRefreshTokenIfNewOneMissing) { auto inner = std::make_shared(io.get_executor()); auto store = std::make_shared(); auto oauth_client = std::make_shared(io.get_executor()); + mcp::auth::MetadataFetchPolicy oauth_policy; + oauth_policy.allowed_origins.push_back("http://127.0.0.1:" + std::to_string(port)); + oauth_policy.allow_plain_http_loopback = true; + oauth_client->set_metadata_policy(oauth_policy); mcp::auth::OAuthConfig config; config.client_id = "test"; diff --git a/test/auth/auth_oauth_test.cpp b/test/auth/auth_oauth_test.cpp index 81cb8e9..9bd8c36 100644 --- a/test/auth/auth_oauth_test.cpp +++ b/test/auth/auth_oauth_test.cpp @@ -5,17 +5,29 @@ #include +#include #include #include +#include #include #include +#include +#include +#include #include #include #include +#include +#include +#include +#include #include #include +#include #include #include +#include +#include #include namespace { @@ -45,6 +57,17 @@ asio::awaitable run_mock_server(asio::ip::tcp::acceptor& acceptor, Accepto co_await handler(std::move(socket), acceptor); } +/// A policy admitting exactly the plain-http loopback fixture these tests drive directly. A +/// policy-less `OAuthHttpClient` defaults to deny-all, so every test that talks to a mock +/// server without going through `OAuthAuthorizationManager` (which always installs its own policy) +/// must opt in explicitly. +mcp::auth::MetadataFetchPolicy loopback_policy(unsigned short port) { + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins.push_back("http://127.0.0.1:" + std::to_string(port)); + policy.allow_plain_http_loopback = true; + return policy; +} + } // namespace TEST(AuthSha256Test, EmptyString) { @@ -295,6 +318,7 @@ TEST_F(MockTokenServer, ExchangeCodeProducesValidToken) { io_ctx_, [&]() -> asio::awaitable { mcp::auth::OAuthHttpClient client(io_ctx_.get_executor()); + client.set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthConfig config; config.client_id = "test_client"; config.token_endpoint = "http://127.0.0.1:" + std::to_string(port) + "/token"; @@ -356,6 +380,7 @@ TEST_F(MockTokenServer, ExchangeCodeWithClientSecret) { io_ctx_, [&]() -> asio::awaitable { mcp::auth::OAuthHttpClient client(io_ctx_.get_executor()); + client.set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthConfig config; config.client_id = "test_client"; config.client_secret = "super_secret"; @@ -413,6 +438,7 @@ TEST_F(MockTokenServer, RefreshTokenProducesNewToken) { io_ctx_, [&]() -> asio::awaitable { mcp::auth::OAuthHttpClient client(io_ctx_.get_executor()); + client.set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthConfig config; config.client_id = "test_client"; config.token_endpoint = "http://127.0.0.1:" + std::to_string(port) + "/token"; @@ -465,6 +491,7 @@ TEST_F(MockTokenServer, TokenExchangeErrorThrows) { io_ctx_, [&]() -> asio::awaitable { mcp::auth::OAuthHttpClient client(io_ctx_.get_executor()); + client.set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthConfig config; config.client_id = "test_client"; config.token_endpoint = "http://127.0.0.1:" + std::to_string(port) + "/token"; @@ -483,6 +510,71 @@ TEST_F(MockTokenServer, TokenExchangeErrorThrows) { EXPECT_TRUE(threw); } +TEST_F(MockTokenServer, DoesNotReplayThePostAfterAnAmbiguousMidResponseFailure) { + // The connection dies after the request is fully read but before any response is written: from + // the client's point of view the server may or may not have acted on it, so this POST must never + // be silently replayed within the same exchange_code() call. Redirects are already never + // followed for a POST (src/auth/oauth.cpp: run_post_json) for the same reason; this proves there + // is no other path that re-sends it. + constexpr unsigned short port = 18111; + + asio::ip::tcp::acceptor acceptor(io_ctx_, {asio::ip::make_address("127.0.0.1"), port}); + int requests_seen = 0; + + asio::co_spawn( + io_ctx_, + [&]() -> asio::awaitable { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + beast::tcp_stream stream(std::move(socket)); + beast::flat_buffer buffer; + http::request req; + co_await http::async_read(stream, buffer, req, asio::use_awaitable); + ++requests_seen; + beast::error_code ec; + stream.socket().close(ec); + + // Give a bounded window for a (forbidden) replay to arrive, then cancel the accept so + // the test never hangs waiting for a connection that -- correctly -- never comes. + asio::steady_timer cutoff(io_ctx_); + cutoff.expires_after(std::chrono::milliseconds(300)); + cutoff.async_wait([&](boost::system::error_code) { acceptor.cancel(); }); + + boost::system::error_code accept_error; + (void)co_await acceptor.async_accept( + asio::redirect_error(asio::use_awaitable, accept_error)); + if (!accept_error) { + ++requests_seen; + } + cutoff.cancel(); + }, + asio::detached); + + bool threw = false; + + asio::co_spawn( + io_ctx_, + [&]() -> asio::awaitable { + mcp::auth::OAuthHttpClient client(io_ctx_.get_executor()); + client.set_metadata_policy(loopback_policy(port)); + mcp::auth::OAuthConfig config; + config.client_id = "test_client"; + config.token_endpoint = "http://127.0.0.1:" + std::to_string(port) + "/token"; + config.redirect_uri = "http://localhost/callback"; + + try { + co_await client.exchange_code(config, "a-code", "verifier"); + } catch (const std::exception&) { + threw = true; + } + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_TRUE(threw); + EXPECT_EQ(requests_seen, 1); +} + TEST_F(MockTokenServer, GetJsonReturnsValidJson) { constexpr unsigned short port = 18099; @@ -518,6 +610,7 @@ TEST_F(MockTokenServer, GetJsonReturnsValidJson) { io_ctx_, [&]() -> asio::awaitable { mcp::auth::OAuthHttpClient client(io_ctx_.get_executor()); + client.set_metadata_policy(loopback_policy(port)); result = co_await client.get_json("http://127.0.0.1:" + std::to_string(port) + "/well-known"); }, @@ -560,6 +653,7 @@ TEST_F(MockTokenServer, GetJsonErrorThrows) { io_ctx_, [&]() -> asio::awaitable { mcp::auth::OAuthHttpClient client(io_ctx_.get_executor()); + client.set_metadata_policy(loopback_policy(port)); try { co_await client.get_json("http://127.0.0.1:" + std::to_string(port) + "/nope"); } catch (const std::runtime_error&) { @@ -630,6 +724,36 @@ TEST(AuthDiscoveryMetadataTest, AuthServerMinimalFields) { EXPECT_FALSE(metadata.scopes_supported.has_value()); } +TEST(AuthDiscoveryMetadataTest, AuthServerW5ExtensionFieldsParsedWhenPresent) { + json j = {{"issuer", "https://auth.example.com"}, + {"authorization_endpoint", "https://auth.example.com/authorize"}, + {"token_endpoint", "https://auth.example.com/token"}, + {"authorization_response_iss_parameter_supported", true}, + {"client_id_metadata_document_supported", true}, + {"token_endpoint_auth_methods_supported", {"client_secret_basic", "none"}}}; + + auto metadata = j.get(); + ASSERT_TRUE(metadata.authorization_response_iss_parameter_supported.has_value()); + EXPECT_TRUE(metadata.authorization_response_iss_parameter_supported.value()); + ASSERT_TRUE(metadata.client_id_metadata_document_supported.has_value()); + EXPECT_TRUE(metadata.client_id_metadata_document_supported.value()); + ASSERT_TRUE(metadata.token_endpoint_auth_methods_supported.has_value()); + ASSERT_EQ(metadata.token_endpoint_auth_methods_supported->size(), 2); + EXPECT_EQ(metadata.token_endpoint_auth_methods_supported->at(0), "client_secret_basic"); + EXPECT_EQ(metadata.token_endpoint_auth_methods_supported->at(1), "none"); +} + +TEST(AuthDiscoveryMetadataTest, AuthServerW5ExtensionFieldsAbsentByDefault) { + json j = {{"issuer", "https://auth.example.com"}, + {"authorization_endpoint", "https://auth.example.com/authorize"}, + {"token_endpoint", "https://auth.example.com/token"}}; + + auto metadata = j.get(); + EXPECT_FALSE(metadata.authorization_response_iss_parameter_supported.has_value()); + EXPECT_FALSE(metadata.client_id_metadata_document_supported.has_value()); + EXPECT_FALSE(metadata.token_endpoint_auth_methods_supported.has_value()); +} + TEST(AuthCacheTest, CachedEntryNotExpired) { mcp::auth::CachedEntry entry{42, std::chrono::steady_clock::now() + std::chrono::hours(1)}; EXPECT_FALSE(entry.is_expired()); @@ -682,6 +806,7 @@ TEST_F(DiscoveryTest, DiscoverProtectedResource) { io_ctx_, [&]() -> asio::awaitable { auto http_client = std::make_shared(io_ctx_.get_executor()); + http_client->set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthDiscoveryClient discovery(http_client, std::chrono::seconds(60)); result = co_await discovery.discover_protected_resource("http://127.0.0.1:" + @@ -735,6 +860,7 @@ TEST_F(DiscoveryTest, DiscoverProtectedResourceWithPath) { io_ctx_, [&]() -> asio::awaitable { auto http_client = std::make_shared(io_ctx_.get_executor()); + http_client->set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthDiscoveryClient discovery(http_client, std::chrono::seconds(60)); result = co_await discovery.discover_protected_resource( @@ -787,6 +913,7 @@ TEST_F(DiscoveryTest, DiscoverAuthServer) { io_ctx_, [&]() -> asio::awaitable { auto http_client = std::make_shared(io_ctx_.get_executor()); + http_client->set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthDiscoveryClient discovery(http_client, std::chrono::seconds(60)); result = @@ -863,6 +990,7 @@ TEST_F(DiscoveryTest, DiscoverAuthServerFallsBackToOIDC) { io_ctx_, [&]() -> asio::awaitable { auto http_client = std::make_shared(io_ctx_.get_executor()); + http_client->set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthDiscoveryClient discovery(http_client, std::chrono::seconds(60)); result = @@ -916,6 +1044,7 @@ TEST_F(DiscoveryTest, CacheHitSkipsNetworkCall) { io_ctx_, [&]() -> asio::awaitable { auto http_client = std::make_shared(io_ctx_.get_executor()); + http_client->set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthDiscoveryClient discovery(http_client, std::chrono::seconds(300)); result1 = @@ -971,6 +1100,7 @@ TEST_F(DiscoveryTest, ClearCacheInvalidatesEntries) { io_ctx_, [&]() -> asio::awaitable { auto http_client = std::make_shared(io_ctx_.get_executor()); + http_client->set_metadata_policy(loopback_policy(port)); mcp::auth::OAuthDiscoveryClient discovery(http_client, std::chrono::seconds(300)); co_await discovery.discover_auth_server("http://127.0.0.1:" + std::to_string(port)); @@ -983,3 +1113,418 @@ TEST_F(DiscoveryTest, ClearCacheInvalidatesEntries) { EXPECT_EQ(request_count, 2); } + +// A policy-less client denies every origin: no origin is on the (empty) allow list, and plain http +// is refused outright. Refusal happens in `enforce_url_policy` before any lookup, so the host +// resolver this test installs must never run. +TEST(AuthOAuthHttpClientDefaultPolicyTest, PolicyLessClientRefusesAnHttpLoopbackUrlWithoutResolving) { + asio::io_context io_ctx; + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + + int resolver_calls = 0; + client.set_host_resolver( + [&resolver_calls](const std::string&, const std::string&) -> std::vector { + ++resolver_calls; + return {"127.0.0.1"}; + }); + + bool threw = false; + auto decision = mcp::auth::MetadataUrlDecision::allowed; + + asio::co_spawn( + io_ctx, + [&]() -> asio::awaitable { + try { + (void)co_await client.get_json("http://127.0.0.1:18199/probe"); + } catch (const mcp::auth::MetadataPolicyError& error) { + threw = true; + decision = error.decision(); + } + }, + asio::detached); + + io_ctx.run(); + + EXPECT_TRUE(threw); + EXPECT_TRUE(decision == mcp::auth::MetadataUrlDecision::origin_not_allowed || + decision == mcp::auth::MetadataUrlDecision::scheme_not_allowed) + << "unexpected decision: " << mcp::auth::describe(decision); + EXPECT_EQ(resolver_calls, 0); +} + +// ------------------------------------------------------------------------------------------- +// The metadata policy and the host resolver must be installable as ONE change. +// +// The fixture is two HTTP servers sharing a port on two loopback addresses, each naming itself +// in its body. Every request targets `http://localhost:/doc`, so the installed resolver +// alone decides which server answers, and the body IS the classification: a body of "new" can +// only be produced by the NEW resolver under the OLD, wider allow list -- the defect, on the wire. +// ------------------------------------------------------------------------------------------- +namespace { + +/// One loopback HTTP server that answers every request with its own name. +class NamedLoopbackServer { + public: + NamedLoopbackServer(asio::io_context& ctx, const std::string& address, unsigned short port, + std::string name) + : acceptor_(ctx, {asio::ip::make_address(address), port}), name_(std::move(name)) { + asio::co_spawn(ctx, accept_loop(), asio::detached); + } + + void close() { + boost::system::error_code ec; + acceptor_.close(ec); + } + + private: + asio::awaitable accept_loop() { + for (;;) { + boost::system::error_code ec; + auto socket = + co_await acceptor_.async_accept(asio::redirect_error(asio::use_awaitable, ec)); + if (ec) { + co_return; + } + asio::co_spawn(socket.get_executor(), serve(std::move(socket)), asio::detached); + } + } + + asio::awaitable serve(asio::ip::tcp::socket socket) { + beast::tcp_stream stream(std::move(socket)); + beast::flat_buffer buffer; + http::request request; + boost::system::error_code ec; + co_await http::async_read(stream, buffer, request, + asio::redirect_error(asio::use_awaitable, ec)); + if (ec) { + co_return; + } + http::response response{http::status::ok, request.version()}; + response.set(http::field::content_type, "application/json"); + response.body() = json{{"server", name_}}.dump(); + response.prepare_payload(); + co_await http::async_write(stream, response, asio::redirect_error(asio::use_awaitable, ec)); + stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, ec); + } + + asio::ip::tcp::acceptor acceptor_; + std::string name_; +}; + +/// Whether this environment has promised a second loopback address. +/// +/// Every OAuthSetterPairAtomicity test needs a port free on both 127.0.0.1 and 127.0.0.2, and skips +/// when there is none. On a host or container without 127.0.0.2 that silently skips every one of them +/// while the suite still reports green. Where the second address is known to exist (Linux, where +/// 127.0.0.0/8 is bound whole), set MCP_REQUIRE_TWIN_LOOPBACK=1 and a missing port fails the run +/// instead of skipping it. CI sets it on Linux; the default stays a skip so the suite remains +/// runnable anywhere. +[[nodiscard]] bool twin_loopback_is_required() { + const char* const flag = std::getenv("MCP_REQUIRE_TWIN_LOOPBACK"); + return flag != nullptr && std::string_view(flag) == "1"; +} + +mcp::auth::MetadataFetchPolicy origin_policy(const std::string& origin) { + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins.push_back(origin); + policy.allow_plain_http_loopback = true; + return policy; +} + +/// What one completed exchange was observed to do. +enum class Observed { + old_server, + new_server, + refused_origin, + other +}; + +Observed classify_exchange(const std::exception_ptr& failure, const json& body) { + if (failure == nullptr) { + const auto name = body.value("server", std::string{}); + if (name == "old") { + return Observed::old_server; + } + if (name == "new") { + return Observed::new_server; + } + return Observed::other; + } + try { + std::rethrow_exception(failure); + } catch (const mcp::auth::MetadataPolicyError& error) { + return error.decision() == mcp::auth::MetadataUrlDecision::origin_not_allowed + ? Observed::refused_origin + : Observed::other; + } catch (const std::exception&) { + return Observed::other; + } +} + +void busy_wait_micros(int micros) { + const auto deadline = std::chrono::steady_clock::now() + std::chrono::microseconds(micros); + while (std::chrono::steady_clock::now() < deadline) { + } +} + +/// The two servers plus the two configurations the application moves between. +struct TwinFixture { + explicit TwinFixture(unsigned short port) + : old_server(ctx, "127.0.0.1", port, "old"), + new_server(ctx, "127.0.0.2", port, "new"), + url("http://localhost:" + std::to_string(port) + "/doc"), + wide(origin_policy("http://localhost:" + std::to_string(port))), + narrow(origin_policy("http://127.0.0.9:1")) { + thread = std::thread([this]() { ctx.run(); }); + } + + ~TwinFixture() { + asio::post(ctx, [this]() { + old_server.close(); + new_server.close(); + }); + work.reset(); + if (thread.joinable()) { + thread.join(); + } + } + + asio::io_context ctx; + asio::executor_work_guard work{asio::make_work_guard(ctx)}; + NamedLoopbackServer old_server; + NamedLoopbackServer new_server; + std::thread thread; + std::string url; + mcp::auth::MetadataFetchPolicy wide; + mcp::auth::MetadataFetchPolicy narrow; + mcp::auth::HostResolver resolver_old = [](const std::string&, const std::string&) { + return std::vector{"127.0.0.1"}; + }; + mcp::auth::HostResolver resolver_new = [](const std::string&, const std::string&) { + return std::vector{"127.0.0.2"}; + }; +}; + +/// The fixture on the first port free on BOTH loopback addresses, or null when there is none +/// because the second address is unavailable. +/// +/// The fixture binds the port itself and keeps it. Probing for a free port and binding it +/// afterwards leaves a gap, and ctest runs these tests as parallel processes that all search the +/// same range: two of them found the same port free and the slower one failed to bind it. +std::unique_ptr open_twin_fixture() { + for (unsigned short port = 18140; port < 18200; ++port) { + try { + return std::make_unique(port); + } catch (const boost::system::system_error&) { + continue; + } + } + return nullptr; +} + +/// Skip the calling test when no twin loopback port is available -- or fail it, when the +/// environment declared that one must be. Declares `fixture_name` as the fixture to use. +#define MCP_TWIN_FIXTURE_OR_SKIP(fixture_name) \ + const std::unique_ptr fixture_name##_owner = open_twin_fixture(); \ + if (fixture_name##_owner == nullptr) { \ + if (twin_loopback_is_required()) { \ + FAIL() << "MCP_REQUIRE_TWIN_LOOPBACK=1, but no port in [18140, 18200) is free on " \ + "both 127.0.0.1 and 127.0.0.2, so this test would have skipped and the " \ + "setter-pair atomicity evidence would have vanished silently"; \ + } \ + GTEST_SKIP() << "no port free on both 127.0.0.1 and 127.0.0.2"; \ + } \ + TwinFixture& fixture_name = *fixture_name##_owner + +} // namespace + +// The hazard itself, with the race taken out of it: the exchange is started at a point the test +// chooses, inside the gap between the two setter calls. It sees the new resolver under the old, +// wider allow list every time, because that is simply what the client's state is at that instant. +TEST(OAuthSetterPairAtomicity, TheTwoSingleSettersLeaveAWindowAnExchangeCanFallInto) { + MCP_TWIN_FIXTURE_OR_SKIP(fixture); + + asio::io_context client_ctx; + auto client = std::make_shared(client_ctx.get_executor()); + client->set_host_resolver(fixture.resolver_old); + client->set_metadata_policy(fixture.wide); + + std::promise resolver_installed; + std::promise exchange_finished; + auto installed = resolver_installed.get_future(); + auto finished = exchange_finished.get_future(); + + std::thread reconfigurer([&]() { + client->set_host_resolver(fixture.resolver_new); + resolver_installed.set_value(); + finished.wait(); + client->set_metadata_policy(fixture.narrow); + }); + + installed.wait(); + + std::exception_ptr failure; + json body; + asio::co_spawn( + client_ctx, + [&]() -> asio::awaitable { + try { + body = co_await client->get_json(fixture.url); + } catch (...) { + failure = std::current_exception(); + } + }, + asio::detached); + client_ctx.run(); + exchange_finished.set_value(); + reconfigurer.join(); + + EXPECT_EQ(classify_exchange(failure, body), Observed::new_server) << "body was " << body.dump(); +} + +// The regression. An application that moves between two whole configurations with a millisecond +// of its own work between the two calls admits, on nearly every narrowing, an exchange directed +// by the new resolver and validated against the old allow list. Applied as one unit there is no +// instant at which that state exists, so the count is zero rather than small. +TEST(OAuthSetterPairAtomicity, ReconfiguringAsOnePairNeverExposesTheNewResolverUnderTheOldPolicy) { + MCP_TWIN_FIXTURE_OR_SKIP(fixture); + + asio::io_context client_ctx; + auto work = asio::make_work_guard(client_ctx); + auto client = std::make_shared(client_ctx.get_executor()); + + auto reconfigure = [&](const mcp::auth::MetadataFetchPolicy& policy, + const mcp::auth::HostResolver& resolver) { + client->configure(policy, resolver); + }; + + reconfigure(fixture.wide, fixture.resolver_old); + + std::atomic old_server{0}; + std::atomic new_server{0}; + std::atomic refused{0}; + std::atomic other{0}; + std::atomic in_flight{0}; + std::atomic transitions{0}; + std::atomic stop{false}; + + std::vector io_threads; + io_threads.reserve(2); + for (int index = 0; index < 2; ++index) { + io_threads.emplace_back([&client_ctx]() { client_ctx.run(); }); + } + + std::thread writer([&]() { + while (!stop.load(std::memory_order_relaxed)) { + reconfigure(fixture.narrow, fixture.resolver_new); + transitions.fetch_add(1, std::memory_order_relaxed); + busy_wait_micros(1000); + reconfigure(fixture.wide, fixture.resolver_old); + busy_wait_micros(1000); + } + }); + + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(20); + while (std::chrono::steady_clock::now() < deadline) { + if (transitions.load() >= 25 && old_server.load() >= 25 && refused.load() >= 25) { + break; + } + if (in_flight.load(std::memory_order_relaxed) >= 2) { + std::this_thread::yield(); + continue; + } + in_flight.fetch_add(1, std::memory_order_relaxed); + asio::co_spawn( + client_ctx, + [&]() -> asio::awaitable { + std::exception_ptr failure; + json body; + try { + body = co_await client->get_json(fixture.url); + } catch (...) { + failure = std::current_exception(); + } + switch (classify_exchange(failure, body)) { + case Observed::old_server: + old_server.fetch_add(1, std::memory_order_relaxed); + break; + case Observed::new_server: + new_server.fetch_add(1, std::memory_order_relaxed); + break; + case Observed::refused_origin: + refused.fetch_add(1, std::memory_order_relaxed); + break; + case Observed::other: + other.fetch_add(1, std::memory_order_relaxed); + break; + } + in_flight.fetch_sub(1, std::memory_order_relaxed); + }, + asio::detached); + } + + stop.store(true); + writer.join(); + const auto drain_deadline = std::chrono::steady_clock::now() + std::chrono::seconds(10); + while (in_flight.load() > 0 && std::chrono::steady_clock::now() < drain_deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(5)); + } + ASSERT_EQ(in_flight.load(), 0u) << "an exchange never completed; the run proves nothing"; + work.reset(); + client_ctx.stop(); + for (auto& thread : io_threads) { + thread.join(); + } + + // Non-vacuity: the traffic must have straddled BOTH whole configurations, or a zero mixed count + // would only mean the requests all landed in one steady state. + EXPECT_GT(old_server.load(), 0u) << "no exchange ever ran under the wide configuration"; + EXPECT_GT(refused.load(), 0u) << "no exchange ever ran under the narrow configuration"; + EXPECT_GE(transitions.load(), 25u) << "the run did not reach enough narrowings to mean much"; + EXPECT_EQ(new_server.load(), 0u) + << new_server.load() << " of " + << (old_server.load() + new_server.load() + refused.load() + other.load()) + << " exchanges were directed by the NEW resolver while validated against the OLD, wider " + "allow list, across " + << transitions.load() << " narrowings"; +} + +// Guards the regression above from passing for the wrong reason: a configure() that quietly +// dropped its resolver argument would also never produce a "new" body. Each call here installs a +// configuration and the exchange that follows must show BOTH halves of it. +TEST(OAuthSetterPairAtomicity, ConfigureInstallsBothOfItsArguments) { + MCP_TWIN_FIXTURE_OR_SKIP(fixture); + + asio::io_context client_ctx; + mcp::auth::OAuthHttpClient client(client_ctx.get_executor()); + + auto observe = [&]() { + std::exception_ptr failure; + json body; + client_ctx.restart(); + asio::co_spawn( + client_ctx, + [&]() -> asio::awaitable { + try { + body = co_await client.get_json(fixture.url); + } catch (...) { + failure = std::current_exception(); + } + }, + asio::detached); + client_ctx.run(); + return classify_exchange(failure, body); + }; + + client.configure(fixture.wide, fixture.resolver_old); + EXPECT_EQ(observe(), Observed::old_server); + + // Only the resolver changes: the policy half must still be the wide one, or this would be a + // refusal rather than a body from the second server. + client.configure(fixture.wide, fixture.resolver_new); + EXPECT_EQ(observe(), Observed::new_server); + + // Only the policy changes: the narrowing must take effect on the very next exchange. + client.configure(fixture.narrow, fixture.resolver_new); + EXPECT_EQ(observe(), Observed::refused_origin); +} diff --git a/test/auth/auth_token_endpoint_auth_test.cpp b/test/auth/auth_token_endpoint_auth_test.cpp new file mode 100644 index 0000000..66d1545 --- /dev/null +++ b/test/auth/auth_token_endpoint_auth_test.cpp @@ -0,0 +1,301 @@ +/** + * @file auth_token_endpoint_auth_test.cpp + * @brief Wire-level and negative-log evidence for token endpoint client authentication: + * the client secret goes exactly where the negotiated `token_endpoint_auth_method` says it + * belongs -- HTTP Basic header for `client_secret_basic`, the form body for + * `client_secret_post`, nowhere at all for `none` -- and it never surfaces in a diagnostic + * message when the token request fails. + * + * client identity selection and the decision function `select_token_endpoint_auth_method` + * already have coverage elsewhere; this file exercises `OAuthHttpClient::exchange_code` directly + * against a loopback token endpoint, one exchange per test, so each assertion is about exactly where + * bytes travelled on the wire. + */ + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace asio = boost::asio; +namespace beast = boost::beast; +namespace http = beast::http; +using json = nlohmann::json; + +namespace { + +/// What a single served request looked like, as far as the client authentication contract cares. +struct RecordedRequest { + std::string target; + std::string body; + std::string authorization; +}; + +/// The outcome of one `exchange_code()` call against a loopback server that serves exactly one +/// request. +struct ExchangeOutcome { + mcp::auth::TokenResponse token; + std::exception_ptr failure; + RecordedRequest recorded; +}; + +http::response json_response(const json& body, + http::status status = http::status::ok) { + http::response response{status, 11}; + response.set(http::field::content_type, "application/json"); + response.body() = body.dump(); + return response; +} + +json token_document() { + return {{"access_token", "granted-access-token"}, {"token_type", "Bearer"}, {"expires_in", 3600}}; +} + +/// Accept one connection, read the request, record it, and reply with `reply`. +asio::awaitable capture_one_request(asio::ip::tcp::acceptor& acceptor, + http::response reply) { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + beast::tcp_stream stream(std::move(socket)); + + beast::flat_buffer buffer; + http::request request; + co_await http::async_read(stream, buffer, request, asio::use_awaitable); + + RecordedRequest recorded; + recorded.target = std::string(request.target()); + recorded.body = request.body(); + recorded.authorization = std::string(request[http::field::authorization]); + + reply.version(request.version()); + reply.prepare_payload(); + co_await http::async_write(stream, reply, asio::use_awaitable); + + beast::error_code shutdown_error; + (void)stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + co_return recorded; +} + +/// Run one `exchange_code()` call against a fresh loopback acceptor that replies with `reply`. +/// `make_config` receives the ephemeral port so it can point `token_endpoint` at it. +ExchangeOutcome run_exchange(const std::function& make_config, + http::response reply) { + ExchangeOutcome outcome; + asio::io_context io_ctx; + asio::ip::tcp::acceptor acceptor(io_ctx, {asio::ip::make_address("127.0.0.1"), 0}); + const auto port = acceptor.local_endpoint().port(); + + asio::co_spawn( + io_ctx, + [&]() -> asio::awaitable { + outcome.recorded = co_await capture_one_request(acceptor, std::move(reply)); + }, + asio::detached); + + asio::co_spawn( + io_ctx, + [&]() -> asio::awaitable { + mcp::auth::OAuthHttpClient client(io_ctx.get_executor()); + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins.push_back("http://127.0.0.1:" + std::to_string(port)); + policy.allow_plain_http_loopback = true; + client.set_metadata_policy(policy); + const auto config = make_config(port); + try { + outcome.token = co_await client.exchange_code(config, "test-authorization-code", + "test-code-verifier"); + } catch (...) { + outcome.failure = std::current_exception(); + } + }, + asio::detached); + + io_ctx.run(); + return outcome; +} + +std::string token_endpoint_url(unsigned short port) { + return "http://127.0.0.1:" + std::to_string(port) + "/token"; +} + +/// What `apply_client_authentication()` in src/auth/oauth.cpp builds for `client_secret_basic`, +/// reproduced here so the expected header is derived rather than hard-coded. +std::string expected_basic_header(const std::string& client_id, const std::string& client_secret) { + const auto credentials = + mcp::auth::detail::url_encode(client_id) + ":" + mcp::auth::detail::url_encode(client_secret); + return "Basic " + + mcp::auth::detail::base64_encode(reinterpret_cast(credentials.data()), + credentials.size()); +} + +/// Unwrap an `std::exception_ptr` captured from the exchange into its `what()` text. +std::string message_of(const std::exception_ptr& failure) { + if (!failure) { + return {}; + } + try { + std::rethrow_exception(failure); + } catch (const std::exception& error) { + return error.what(); + } + return {}; +} + +} // namespace + +TEST(AuthTokenEndpointAuthWireTest, ClientSecretBasicPutsTheSecretOnlyInTheAuthorizationHeader) { + const std::string client_id = "conf-client"; + const std::string client_secret = "sekrit-basic-value"; + + const auto outcome = run_exchange( + [&](unsigned short port) { + mcp::auth::OAuthConfig config; + config.client_id = client_id; + config.client_secret = client_secret; + config.token_endpoint = token_endpoint_url(port); + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.token_endpoint_auth_method = "client_secret_basic"; + return config; + }, + json_response(token_document())); + + ASSERT_EQ(outcome.failure, nullptr); + EXPECT_EQ(outcome.token.access_token, "granted-access-token"); + + // The secret travels in the Authorization header, correctly base64-encoded as `id:secret`. + EXPECT_EQ(outcome.recorded.authorization, expected_basic_header(client_id, client_secret)); + + // It appears nowhere else: not in the form body, not as a `client_secret` parameter, not in the + // request target (path/query). + EXPECT_EQ(outcome.recorded.body.find(client_secret), std::string::npos); + EXPECT_EQ(outcome.recorded.body.find("client_secret"), std::string::npos); + EXPECT_EQ(outcome.recorded.target.find(client_secret), std::string::npos); +} + +TEST(AuthTokenEndpointAuthWireTest, ClientSecretPostPutsTheSecretOnlyInTheFormBody) { + const std::string client_id = "conf-client"; + const std::string client_secret = "sekrit-post-value"; + + const auto outcome = run_exchange( + [&](unsigned short port) { + mcp::auth::OAuthConfig config; + config.client_id = client_id; + config.client_secret = client_secret; + config.token_endpoint = token_endpoint_url(port); + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.token_endpoint_auth_method = "client_secret_post"; + return config; + }, + json_response(token_document())); + + ASSERT_EQ(outcome.failure, nullptr); + EXPECT_EQ(outcome.token.access_token, "granted-access-token"); + + // No Authorization header at all: the secret did not also leak into HTTP Basic. + EXPECT_TRUE(outcome.recorded.authorization.empty()); + + // The secret is present exactly once, as the `client_secret` form parameter. + EXPECT_NE( + outcome.recorded.body.find("client_secret=" + mcp::auth::detail::url_encode(client_secret)), + std::string::npos); + EXPECT_EQ(outcome.recorded.target.find(client_secret), std::string::npos); +} + +TEST(AuthTokenEndpointAuthWireTest, NoneMethodSendsNoSecretEvenWhenOneIsKnown) { + const std::string client_id = "conf-client"; + // A secret the application happens to know (e.g. one a dynamic registration returned earlier), + // but the negotiated method is `none`, so it must never be sent anywhere. + const std::string client_secret = "dcr-issued-secret-that-must-not-travel"; + + const auto outcome = run_exchange( + [&](unsigned short port) { + mcp::auth::OAuthConfig config; + config.client_id = client_id; + config.client_secret = client_secret; + config.token_endpoint = token_endpoint_url(port); + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.token_endpoint_auth_method = "none"; + return config; + }, + json_response(token_document())); + + ASSERT_EQ(outcome.failure, nullptr); + EXPECT_EQ(outcome.token.access_token, "granted-access-token"); + + EXPECT_TRUE(outcome.recorded.authorization.empty()); + EXPECT_EQ(outcome.recorded.body.find("client_secret"), std::string::npos); + EXPECT_EQ(outcome.recorded.body.find(client_secret), std::string::npos); + EXPECT_EQ(outcome.recorded.target.find(client_secret), std::string::npos); +} + +TEST(AuthTokenEndpointAuthNegativeLogTest, + FailingTokenRequestWithClientSecretPostNeverEchoesTheSecretInTheExceptionMessage) { + const std::string client_id = "conf-client"; + const std::string client_secret = "sekrit-post-value-that-must-stay-out-of-logs"; + + // A realistic authorization-server error body: it names the problem, it does not echo the + // request back. + const auto outcome = run_exchange( + [&](unsigned short port) { + mcp::auth::OAuthConfig config; + config.client_id = client_id; + config.client_secret = client_secret; + config.token_endpoint = token_endpoint_url(port); + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.token_endpoint_auth_method = "client_secret_post"; + return config; + }, + json_response( + json{{"error", "invalid_client"}, {"error_description", "client authentication failed"}}, + http::status::unauthorized)); + + ASSERT_NE(outcome.failure, nullptr); + // The secret did travel in this request's body (client_secret_post); the assertion that matters + // is that the *diagnostic surfaced to the caller* -- the thrown exception's message -- never + // repeats it. + ASSERT_NE(outcome.recorded.body.find(client_secret), std::string::npos) + << "test setup error: the secret should have been in the request body for this scenario"; + + const auto message = message_of(outcome.failure); + EXPECT_FALSE(message.empty()); + EXPECT_EQ(message.find(client_secret), std::string::npos) + << "exception message leaked the client secret: " << message; +} + +TEST(AuthTokenEndpointAuthNegativeLogTest, + FailingTokenRequestWithClientSecretBasicNeverEchoesTheSecretInTheExceptionMessage) { + const std::string client_id = "conf-client"; + const std::string client_secret = "sekrit-basic-value-that-must-stay-out-of-logs"; + + const auto outcome = run_exchange( + [&](unsigned short port) { + mcp::auth::OAuthConfig config; + config.client_id = client_id; + config.client_secret = client_secret; + config.token_endpoint = token_endpoint_url(port); + config.redirect_uri = "http://127.0.0.1:9999/callback"; + config.token_endpoint_auth_method = "client_secret_basic"; + return config; + }, + json_response(json{{"error", "invalid_client"}}, http::status::unauthorized)); + + ASSERT_NE(outcome.failure, nullptr); + ASSERT_EQ(outcome.recorded.authorization, expected_basic_header(client_id, client_secret)) + << "test setup error: the secret should have travelled in the Authorization header"; + + const auto message = message_of(outcome.failure); + EXPECT_FALSE(message.empty()); + EXPECT_EQ(message.find(client_secret), std::string::npos) + << "exception message leaked the client secret: " << message; + // The raw Basic-auth header value must not appear either, since decoding it recovers the secret. + EXPECT_EQ(message.find(expected_basic_header(client_id, client_secret)), std::string::npos); +} diff --git a/test/client/client_core_test.cpp b/test/client/client_core_test.cpp index 5d1a82c..18ddb3e 100644 --- a/test/client/client_core_test.cpp +++ b/test/client/client_core_test.cpp @@ -1,16 +1,27 @@ #include "../test_utils.hpp" #include "mcp/client/client.hpp" +#include "mcp/transport/memory.hpp" #include +#include #include #include +#include #include +#include +#include +#include +#include +#include #include #include +#include #include #include +#include +#include namespace { @@ -20,6 +31,249 @@ nlohmann::json make_initialize_result() { {"serverInfo", {{"name", "test-server"}, {"version", "1.0"}}}}; } +class ImmediateResponseTransport final : public mcp::ITransport { + public: + explicit ImmediateResponseTransport(const boost::asio::any_io_executor& executor) + : executor_(executor), signal_(executor) { + signal_.expires_at(std::chrono::steady_clock::time_point::max()); + } + + mcp::Task read_message() override { + for (;;) { + if (!incoming_.empty()) { + auto message = std::move(incoming_.front()); + incoming_.pop(); + co_return message; + } + if (closed_) { + throw std::runtime_error("transport closed"); + } + + signal_.expires_at(std::chrono::steady_clock::time_point::max()); + try { + co_await signal_.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + } + } + + mcp::Task write_message(std::string_view message) override { + const auto request = nlohmann::json::parse(message); + if (request.contains("id")) { + const auto id = request.at("id").get(); + if (request.value("method", "") == "initialize") { + incoming_.push(make_result_response(id, make_initialize_result()).dump()); + } else { + incoming_.push(make_result_response(id, {{"ok", true}}).dump()); + } + signal_.cancel(); + + // Let the read loop consume and dispatch the response before this + // write completes. This reproduces completion-before-wait ordering. + co_await boost::asio::post(executor_, boost::asio::use_awaitable); + co_await boost::asio::post(executor_, boost::asio::use_awaitable); + } + } + + void close() override { + closed_ = true; + signal_.cancel(); + } + + private: + boost::asio::any_io_executor executor_; + boost::asio::steady_timer signal_; + std::queue incoming_; + bool closed_{false}; +}; + +class BlockingRequestWriteTransport final : public mcp::ITransport { + public: + explicit BlockingRequestWriteTransport(const boost::asio::any_io_executor& executor) + : read_signal_(executor), write_signal_(executor) { + read_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + write_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + } + + mcp::Task read_message() override { + for (;;) { + if (!incoming_.empty()) { + auto message = std::move(incoming_.front()); + incoming_.pop(); + co_return message; + } + if (closed_) { + throw std::runtime_error("transport closed"); + } + + read_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + try { + co_await read_signal_.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + } + } + + mcp::Task write_message(std::string_view message) override { + const auto request = nlohmann::json::parse(message); + if (request.value("method", "") == "initialize") { + incoming_.push( + make_result_response(request.at("id").get(), make_initialize_result()) + .dump()); + read_signal_.cancel(); + co_return; + } + if (!request.contains("id")) { + co_return; + } + + blocked_write_started_.store(true, std::memory_order_release); + write_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + try { + co_await write_signal_.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + throw std::runtime_error("blocked write cancelled"); + } + + void close() override { + closed_ = true; + read_signal_.cancel(); + write_signal_.cancel(); + } + + [[nodiscard]] bool blocked_write_started() const { + return blocked_write_started_.load(std::memory_order_acquire); + } + + private: + boost::asio::steady_timer read_signal_; + boost::asio::steady_timer write_signal_; + std::queue incoming_; + std::atomic_bool blocked_write_started_{false}; + bool closed_{false}; +}; + +class QueuedWriteTransport final : public mcp::ITransport { + public: + explicit QueuedWriteTransport(const boost::asio::any_io_executor& executor) + : read_signal_(executor), blocked_signal_(executor), release_signal_(executor) { + read_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + blocked_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + release_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + } + + mcp::Task read_message() override { + for (;;) { + if (!incoming_.empty()) { + auto message = std::move(incoming_.front()); + incoming_.pop(); + co_return message; + } + if (closed_) { + throw std::runtime_error("transport closed"); + } + + read_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + try { + co_await read_signal_.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + } + } + + mcp::Task write_message(std::string_view message) override { + const auto request = nlohmann::json::parse(message); + const auto method = request.value("method", ""); + methods_.push_back(method); + + if (method == "initialize") { + incoming_.push( + make_result_response(request.at("id").get(), make_initialize_result()) + .dump()); + read_signal_.cancel(); + co_return; + } + if (method == "notifications/block") { + blocked_.store(true, std::memory_order_release); + blocked_signal_.cancel(); + release_signal_.expires_at(std::chrono::steady_clock::time_point::max()); + try { + co_await release_signal_.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + if (closed_) { + throw std::runtime_error("transport closed"); + } + co_return; + } + if (request.contains("id")) { + incoming_.push( + make_result_response(request.at("id").get(), nlohmann::json::object()) + .dump()); + read_signal_.cancel(); + } + } + + mcp::Task wait_until_blocked() { + while (!blocked_.load(std::memory_order_acquire)) { + try { + co_await blocked_signal_.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + } + } + + void release_blocked_write() { release_signal_.cancel(); } + + void close() override { + closed_ = true; + read_signal_.cancel(); + blocked_signal_.cancel(); + release_signal_.cancel(); + } + + [[nodiscard]] const std::vector& methods() const { return methods_; } + + private: + boost::asio::steady_timer read_signal_; + boost::asio::steady_timer blocked_signal_; + boost::asio::steady_timer release_signal_; + std::queue incoming_; + std::vector methods_; + std::atomic_bool blocked_{false}; + bool closed_{false}; +}; + +void run_on_thread_pool(boost::asio::io_context& io_context, std::size_t thread_count = 4) { + std::vector threads; + threads.reserve(thread_count); + for (std::size_t index = 0; index < thread_count; ++index) { + threads.emplace_back([&io_context]() { io_context.run(); }); + } + for (auto& thread : threads) { + thread.join(); + } +} + } // namespace class ClientCoreTest : public ::testing::Test { @@ -68,6 +322,16 @@ TEST_F(ClientCoreTest, SendRequestAddsToMapAndReturnsResult) { EXPECT_EQ(client.pending_request_count(), 0); } +TEST_F(ClientCoreTest, RequestsFailImmediatelyBeforeConnect) { + auto transport = std::make_shared(io_ctx_.get_executor()); + mcp::Client client(transport, io_ctx_.get_executor()); + + EXPECT_THROW(static_cast(client.send_request("ping", std::nullopt)), mcp::McpError); + EXPECT_THROW(static_cast(client.send_notification("notifications/test", std::nullopt)), + mcp::McpError); + EXPECT_THROW(static_cast(client.ping()), mcp::McpError); +} + TEST_F(ClientCoreTest, ConnectHandshakeFollowsMcpOrder) { auto transport = std::make_shared(io_ctx_.get_executor()); auto* raw_transport = transport.get(); @@ -116,6 +380,41 @@ TEST_F(ClientCoreTest, ConnectHandshakeFollowsMcpOrder) { EXPECT_EQ(init_result.serverInfo.version, "1.0"); } +TEST_F(ClientCoreTest, NullIdErrorDoesNotStopReadLoop) { + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + + raw_transport->set_on_write([raw_transport](std::string_view msg) { + const auto json_msg = nlohmann::json::parse(msg); + if (!json_msg.contains("id")) { + return; + } + + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, + {"id", nullptr}, + {"error", {{"code", mcp::g_PARSE_ERROR}, {"message", "parse error"}}}} + .dump()); + raw_transport->enqueue_message( + make_result_response(json_msg.at("id").get(), make_initialize_result()) + .dump()); + }); + + mcp::Client client(transport, io_ctx_.get_executor()); + mcp::InitializeResult result; + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + result = co_await client.connect("test-client", "0.1"); + raw_transport->close(); + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(result.serverInfo.name, "test-server"); +} + TEST_F(ClientCoreTest, SendRequestErrorResponseThrows) { auto transport = std::make_shared(io_ctx_.get_executor()); auto* raw_transport = transport.get(); @@ -421,3 +720,752 @@ TEST_F(ClientCoreTest, PingSendsRequestAndReceivesResponse) { EXPECT_TRUE(ping_completed); EXPECT_EQ(client.pending_request_count(), 0); } + +TEST_F(ClientCoreTest, RequestTimeoutThrowsStructuredMcpError) { + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + + raw_transport->set_on_write([raw_transport](std::string_view msg) { + auto json_msg = nlohmann::json::parse(msg); + if (json_msg.contains("id") && json_msg.value("method", "") == "initialize") { + auto id = json_msg["id"].get(); + raw_transport->enqueue_message(make_result_response(id, make_initialize_result()).dump()); + } + }); + + mcp::Client client(transport, io_ctx_.get_executor()); + std::optional error_code; + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await client.connect("test-client", "0.1"); + try { + mcp::RequestOptions options; + options.timeout = std::chrono::milliseconds(5); + co_await client.send_request("tools/list", std::nullopt, options); + } catch (const mcp::McpError& error) { + error_code = error.code(); + } + client.close(); + }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_TRUE(error_code.has_value()); + EXPECT_EQ(*error_code, mcp::g_REQUEST_TIMEOUT); + EXPECT_EQ(client.pending_request_count(), 0); +} + +TEST_F(ClientCoreTest, RequestTimeoutIncludesTimeQueuedInTransportWrite) { + auto transport = std::make_shared(io_ctx_.get_executor()); + mcp::Client client(transport, io_ctx_.get_executor()); + + std::optional error_code; + bool request_completed = false; + bool watchdog_fired = false; + boost::asio::steady_timer watchdog(io_ctx_); + watchdog.expires_after(std::chrono::milliseconds(250)); + watchdog.async_wait([&](const boost::system::error_code& error) { + if (!error) { + watchdog_fired = true; + client.close(); + } + }); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await client.connect("test-client", "0.1"); + try { + mcp::RequestOptions options; + options.timeout = std::chrono::milliseconds(10); + co_await client.send_request("tools/list", std::nullopt, options); + } catch (const mcp::McpError& error) { + error_code = error.code(); + } + request_completed = true; + watchdog.cancel(); + client.close(); + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_TRUE(transport->blocked_write_started()); + EXPECT_TRUE(request_completed); + EXPECT_FALSE(watchdog_fired); + ASSERT_TRUE(error_code.has_value()); + EXPECT_EQ(*error_code, mcp::g_REQUEST_TIMEOUT); + EXPECT_EQ(client.pending_request_count(), 0); +} + +TEST_F(ClientCoreTest, TimedOutQueuedRequestIsNotTransmittedLater) { + auto transport = std::make_shared(io_ctx_.get_executor()); + mcp::Client client(transport, io_ctx_.get_executor()); + + std::optional error_code; + std::exception_ptr notification_error; + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await client.connect("test-client", "0.1"); + boost::asio::co_spawn(io_ctx_, + client.send_notification("notifications/block", std::nullopt), + [¬ification_error](std::exception_ptr error) { + notification_error = std::move(error); + }); + co_await transport->wait_until_blocked(); + + try { + mcp::RequestOptions options; + options.timeout = std::chrono::milliseconds(10); + co_await client.send_request("tools/list", std::nullopt, options); + } catch (const mcp::McpError& error) { + error_code = error.code(); + } + + transport->release_blocked_write(); + boost::asio::steady_timer settle(io_ctx_, std::chrono::milliseconds(10)); + co_await settle.async_wait(boost::asio::use_awaitable); + client.close(); + }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_TRUE(error_code.has_value()); + EXPECT_EQ(*error_code, mcp::g_REQUEST_TIMEOUT); + EXPECT_EQ(notification_error, nullptr); + EXPECT_EQ(std::count(transport->methods().begin(), transport->methods().end(), "tools/list"), 0); + EXPECT_EQ(client.pending_request_count(), 0); +} + +TEST_F(ClientCoreTest, RemoteErrorPreservesJsonRpcData) { + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + + raw_transport->set_on_write([raw_transport](std::string_view msg) { + auto json_msg = nlohmann::json::parse(msg); + if (!json_msg.contains("id")) { + return; + } + auto id = json_msg["id"].get(); + if (json_msg.value("method", "") == "initialize") { + raw_transport->enqueue_message(make_result_response(id, make_initialize_result()).dump()); + return; + } + auto response = make_error_response(id, mcp::g_INVALID_PARAMS, "invalid arguments"); + response["error"]["data"] = {{"field", "limit"}}; + raw_transport->enqueue_message(response.dump()); + }); + + mcp::Client client(transport, io_ctx_.get_executor()); + std::optional captured_error; + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await client.connect("test-client", "0.1"); + try { + co_await client.send_request("tools/call", nlohmann::json::object()); + } catch (const mcp::McpError& error) { + captured_error = error.error(); + } + client.close(); + }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_TRUE(captured_error.has_value()); + EXPECT_EQ(captured_error->code, mcp::g_INVALID_PARAMS); + ASSERT_TRUE(captured_error->data.has_value()); + EXPECT_EQ(captured_error->data->at("field"), "limit"); +} + +TEST_F(ClientCoreTest, ImmediateResponseCannotBeLostBeforeWaitStarts) { + auto transport = std::make_shared(io_ctx_.get_executor()); + mcp::ClientOptions client_options; + client_options.request_timeout = std::chrono::seconds(2); + mcp::Client client(transport, io_ctx_.get_executor(), client_options); + + boost::asio::steady_timer watchdog(io_ctx_); + watchdog.expires_after(std::chrono::milliseconds(250)); + bool watchdog_fired = false; + bool request_completed = false; + nlohmann::json result; + + watchdog.async_wait([&](const boost::system::error_code& error) { + if (!error) { + watchdog_fired = true; + client.close(); + } + }); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await client.connect("immediate-response-client", "1.0"); + result = co_await client.send_request("test/immediate", std::nullopt); + request_completed = true; + watchdog.cancel(); + client.close(); + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_FALSE(watchdog_fired); + EXPECT_TRUE(request_completed); + EXPECT_EQ(result, nlohmann::json({{"ok", true}})); +} + +TEST_F(ClientCoreTest, DestructionSafelyStopsActiveReadLoop) { + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + raw_transport->set_on_write([raw_transport](std::string_view message) { + const auto request = nlohmann::json::parse(message); + if (request.contains("id") && request.value("method", "") == "initialize") { + const auto id = request.at("id").get(); + raw_transport->enqueue_message(make_result_response(id, make_initialize_result()).dump()); + } + }); + + auto client = std::make_unique(transport, io_ctx_.get_executor()); + bool initialized = false; + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await client->connect("destruction-client", "1.0"); + initialized = true; + client.reset(); + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_TRUE(initialized); + EXPECT_EQ(client, nullptr); + EXPECT_TRUE(raw_transport->is_closed()); +} + +TEST_F(ClientCoreTest, ConcurrentRequestsRemainCorrelatedOnMultiThreadedExecutor) { + using namespace std::chrono_literals; + + constexpr std::size_t request_count = 256; + auto [client_transport, server_transport] = + mcp::create_memory_transport_pair(io_ctx_.get_executor()); + mcp::ClientOptions options; + // This timeout and the watchdog below only end a run that has gone wrong. They are wide + // because a sanitizer can stop every thread for seconds while it prepares a report, including + // one a suppression then discards. + options.request_timeout = 30s; + mcp::Client client(client_transport, io_ctx_.get_executor(), options); + + std::atomic_size_t completed{0}; + std::atomic_size_t correct{0}; + std::atomic_size_t errors{0}; + std::atomic_bool stop_poll{false}; + + boost::asio::co_spawn( + io_ctx_, + [server_transport]() -> mcp::Task { + std::size_t responses = 0; + while (responses <= request_count) { + const auto request = nlohmann::json::parse(co_await server_transport->read_message()); + if (!request.contains("id")) { + continue; + } + + nlohmann::json result; + if (request.value("method", "") == "initialize") { + result = make_initialize_result(); + } else { + result = {{"sequence", request.at("params").at("sequence")}}; + } + co_await server_transport->write_message( + make_result_response(request.at("id").get(), std::move(result)) + .dump()); + ++responses; + } + }, + [&errors](std::exception_ptr error) { + if (error) { + errors.fetch_add(1, std::memory_order_relaxed); + } + }); + + boost::asio::co_spawn( + io_ctx_, + [&client, &completed, &correct, &errors]() -> mcp::Task { + try { + co_await client.connect("threaded-client", "1.0"); + auto executor = co_await boost::asio::this_coro::executor; + for (std::size_t sequence = 0; sequence < request_count; ++sequence) { + boost::asio::co_spawn( + executor, + client.send_request("test/repeat", nlohmann::json{{"sequence", sequence}}), + [sequence, &completed, &correct, &errors](std::exception_ptr error, + nlohmann::json result) { + if (error) { + errors.fetch_add(1, std::memory_order_relaxed); + } else if (result.at("sequence").get() == sequence) { + correct.fetch_add(1, std::memory_order_relaxed); + } + completed.fetch_add(1, std::memory_order_release); + }); + } + } catch (...) { + errors.fetch_add(1, std::memory_order_relaxed); + client.close(); + } + }, + boost::asio::detached); + + boost::asio::steady_timer watchdog(io_ctx_); + watchdog.expires_after(60s); + watchdog.async_wait([&client, &completed, &stop_poll](const boost::system::error_code& error) { + if (!error && completed.load(std::memory_order_acquire) != request_count) { + stop_poll.store(true, std::memory_order_release); + client.close(); + } + }); + + boost::asio::steady_timer completion_poll(io_ctx_); + std::function poll; + poll = [&]() { + if (stop_poll.load(std::memory_order_acquire)) { + return; + } + if (completed.load(std::memory_order_acquire) == request_count) { + stop_poll.store(true, std::memory_order_release); + watchdog.cancel(); + client.close(); + return; + } + completion_poll.expires_after(1ms); + completion_poll.async_wait([&poll](const boost::system::error_code& error) { + if (!error) { + poll(); + } + }); + }; + poll(); + + run_on_thread_pool(io_ctx_); + + EXPECT_EQ(completed.load(), request_count); + EXPECT_EQ(correct.load(), request_count); + EXPECT_EQ(errors.load(), 0); + EXPECT_EQ(client.pending_request_count(), 0); +} + +namespace { + +// Drives a connect() whose peer answers `initialize` with the given protocolVersion, and returns +// the McpError message the client raised. The version string is entirely the server's choice, so +// this is the untrusted-peer text that reaches the diagnostic. +struct RejectedVersionOutcome { + bool threw{false}; + int code{0}; + std::string message; +}; + +RejectedVersionOutcome connect_against_protocol_version(boost::asio::io_context& io_ctx, + const std::string& protocol_version) { + auto transport = std::make_shared(io_ctx.get_executor()); + auto* raw_transport = transport.get(); + + raw_transport->set_on_write([raw_transport, protocol_version](std::string_view msg) { + const auto json_msg = nlohmann::json::parse(msg); + if (!json_msg.contains("id")) { + return; + } + raw_transport->enqueue_message( + make_result_response( + json_msg.at("id").get(), + nlohmann::json{{"protocolVersion", protocol_version}, + {"capabilities", nlohmann::json::object()}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0"}}}}) + .dump()); + }); + + mcp::Client client(transport, io_ctx.get_executor()); + + RejectedVersionOutcome outcome; + boost::asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::Implementation info; + info.name = "test-client"; + info.version = "0.1"; + try { + static_cast(co_await client.connect(std::move(info), mcp::ClientCapabilities{})); + } catch (const mcp::McpError& error) { + outcome.threw = true; + outcome.code = error.code(); + outcome.message = error.message(); + } + raw_transport->close(); + }, + boost::asio::detached); + + io_ctx.run(); + return outcome; +} + +} // namespace + +// The protocol version in an `initialize` result is chosen by the remote server. When the client +// refuses it, the rejected value is interpolated into the McpError the application logs, so a +// malicious or broken server that answers with CR/LF forges a line in the application's log and a +// bidi override reorders the rest of the message. +TEST_F(ClientCoreTest, UnsupportedProtocolVersionDiagnosticFlattensServerChosenVersion) { + // "zqtripwire" is a token no other code path produces; see the tripwire assertions below. + // \xe2\x80\xae is U+202E RIGHT-TO-LEFT OVERRIDE, written escaped so this source file does not + // itself contain a bidi override. + const std::string forged_version = + "zqtripwire\r\n2026-09-20 ERROR forged line from the server\xe2\x80\xae reordered tail"; + + const auto outcome = connect_against_protocol_version(io_ctx_, forged_version); + + // Tripwire. The rejection must be the strict-protocol-validation site, not a JSON parse + // failure, a transport error, or any other throw that never interpolates the peer's bytes -- + // each of those would make the assertions below pass for a reason unrelated to the site under + // test. Dump the message on failure so a vacuous pass cannot hide. + ASSERT_TRUE(outcome.threw) << "connect() did not reject the version at all"; + ASSERT_EQ(outcome.code, mcp::g_INVALID_REQUEST) << "actual message: " << outcome.message; + ASSERT_EQ(outcome.message.rfind("Server selected unsupported protocol version: ", 0), 0U) + << "actual message: " << outcome.message; + ASSERT_NE(outcome.message.find("zqtripwire"), std::string::npos) + << "actual message: " << outcome.message; + + EXPECT_EQ(outcome.message.find('\r'), std::string::npos) << "actual message: " << outcome.message; + EXPECT_EQ(outcome.message.find('\n'), std::string::npos) << "actual message: " << outcome.message; + EXPECT_EQ(outcome.message.find("\xe2\x80\xae"), std::string::npos) + << "actual message: " << outcome.message; +} + +// The same site must also bound the value, so a server cannot flood the application's log through +// a megabyte-long protocol version. +TEST_F(ClientCoreTest, UnsupportedProtocolVersionDiagnosticBoundsServerChosenVersion) { + const std::string forged_version = "zqtripwire" + std::string(64 * 1024, 'A'); + + const auto outcome = connect_against_protocol_version(io_ctx_, forged_version); + + ASSERT_TRUE(outcome.threw) << "connect() did not reject the version at all"; + ASSERT_EQ(outcome.code, mcp::g_INVALID_REQUEST) + << "actual message prefix: " << outcome.message.substr(0, 80); + ASSERT_EQ(outcome.message.rfind("Server selected unsupported protocol version: ", 0), 0U) + << "actual message prefix: " << outcome.message.substr(0, 80); + ASSERT_NE(outcome.message.find("zqtripwire"), std::string::npos) + << "actual message prefix: " << outcome.message.substr(0, 80); + + EXPECT_LT(outcome.message.size(), forged_version.size()) + << "message size: " << outcome.message.size(); + EXPECT_LE(outcome.message.size(), std::size_t{512}) << "message size: " << outcome.message.size(); +} + +namespace { + +// Drives a request the peer answers with a JSON-RPC error object, and reports what the client +// raised. +// +// This is the shortest path there is from a peer's bytes to an application's log. `error.message` +// is deserialized verbatim off the wire (`json_message.at("error").get()`), carried into +// McpError, and interpolated into what(). No OAuth, no discovery, no metadata document: an +// ordinary error response to an ordinary request. +struct PeerErrorOutcome { + bool threw{false}; + int code{0}; + std::string what; ///< McpError::what() -- the human-readable diagnostic. + std::string raw_message; ///< McpError::message() -- the structured field, which stays raw. +}; + +PeerErrorOutcome send_request_against_peer_error(boost::asio::io_context& io_ctx, + const std::string& peer_message) { + auto transport = std::make_shared(io_ctx.get_executor()); + auto* raw_transport = transport.get(); + + raw_transport->set_on_write([raw_transport, peer_message](std::string_view msg) { + const auto json_msg = nlohmann::json::parse(msg); + if (!json_msg.contains("id")) { + return; + } + const auto id = json_msg.at("id").get(); + if (json_msg.value("method", "") == "initialize") { + raw_transport->enqueue_message(make_result_response(id, make_initialize_result()).dump()); + return; + } + // The message travels as JSON, so CR/LF written here arrive at the client as real control + // bytes rather than as the two-character escapes: the encoder escapes them, the decoder + // turns them back. That decode is what makes a parsed field sharper than raw wire bytes. + raw_transport->enqueue_message( + make_error_response(id, mcp::g_INTERNAL_ERROR, peer_message).dump()); + }); + + mcp::Client client(transport, io_ctx.get_executor()); + + PeerErrorOutcome outcome; + boost::asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + mcp::Implementation info; + info.name = "test-client"; + info.version = "0.1"; + static_cast(co_await client.connect(std::move(info), mcp::ClientCapabilities{})); + try { + static_cast(co_await client.send_request("tools/list", std::nullopt)); + } catch (const mcp::McpError& error) { + outcome.threw = true; + outcome.code = error.code(); + outcome.what = error.what(); + outcome.raw_message = error.message(); + } + raw_transport->close(); + }, + boost::asio::detached); + + io_ctx.run(); + return outcome; +} + +} // namespace + +// The message in a peer's error response is the peer's text, and it reaches what() -- which is +// what an application logs. A server that answers with CR/LF forges a line in that log, and a bidi +// override reorders the tail of the diagnostic so a refusal can be made to read as its opposite. +TEST_F(ClientCoreTest, PeerErrorDiagnosticFlattensThePeerChosenMessage) { + // "zqtripwire" is a token no other code path produces; see the tripwire assertions below. + // \xe2\x80\xae is U+202E RIGHT-TO-LEFT OVERRIDE, written escaped so this source file does not + // itself contain a bidi override. + const std::string forged = + "denied\r\n2026-09-20 INFO zqtripwire operator approved\xe2\x80\xae reordered tail"; + + const auto outcome = send_request_against_peer_error(io_ctx_, forged); + + // Tripwire, ordered before the flattening assertions. The throw has to be the McpError built + // from the peer's own error object -- not a request timeout, not an envelope-validation + // refusal, not a transport failure. Each of those also raises an McpError, with a message this + // SDK authored, and would satisfy everything below without the peer's bytes ever reaching a + // diagnostic. Dump what() on failure so a vacuous pass cannot hide. + ASSERT_TRUE(outcome.threw) << "send_request() raised nothing at all"; + ASSERT_EQ(outcome.code, mcp::g_INTERNAL_ERROR) << "actual what(): " << outcome.what; + ASSERT_EQ(outcome.what.rfind("JSON-RPC error ", 0), 0U) << "actual what(): " << outcome.what; + ASSERT_NE(outcome.what.find("zqtripwire"), std::string::npos) << "actual what(): " << outcome.what; + + EXPECT_EQ(outcome.what.find('\r'), std::string::npos) << "actual what(): " << outcome.what; + EXPECT_EQ(outcome.what.find('\n'), std::string::npos) << "actual what(): " << outcome.what; + EXPECT_EQ(outcome.what.find("\xe2\x80\xae"), std::string::npos) + << "actual what(): " << outcome.what; + + // The split, and it is the point: what() is a diagnostic and is sanitized; error() and the + // message() it exposes are structured protocol data a caller may compare or re-encode, so they + // must come back exactly as the peer sent them. Sanitizing those instead would silently change + // what an application matches on. + EXPECT_EQ(outcome.raw_message, forged); +} + +// The same site must bound the value too, so a peer cannot flood the application's log through a +// megabyte-long error message. +TEST_F(ClientCoreTest, PeerErrorDiagnosticBoundsThePeerChosenMessage) { + const std::string forged = "zqtripwire" + std::string(64 * 1024, 'A'); + + const auto outcome = send_request_against_peer_error(io_ctx_, forged); + + ASSERT_TRUE(outcome.threw) << "send_request() raised nothing at all"; + ASSERT_EQ(outcome.code, mcp::g_INTERNAL_ERROR) + << "actual what() prefix: " << outcome.what.substr(0, 80); + ASSERT_EQ(outcome.what.rfind("JSON-RPC error ", 0), 0U) + << "actual what() prefix: " << outcome.what.substr(0, 80); + ASSERT_NE(outcome.what.find("zqtripwire"), std::string::npos) + << "actual what() prefix: " << outcome.what.substr(0, 80); + + EXPECT_LT(outcome.what.size(), forged.size()) << "what() size: " << outcome.what.size(); + EXPECT_LE(outcome.what.size(), std::size_t{512}) << "what() size: " << outcome.what.size(); + EXPECT_EQ(outcome.raw_message.size(), forged.size()) + << "the structured message was truncated; only the diagnostic may be"; +} + +namespace { + +// Drives two requests over one session. The peer answers the first with a caller-supplied burst of +// messages and the second normally. +// +// The second request is the whole point: it separates a failure confined to one message from a +// failure that took the session down with it. A client that survives a message it cannot use fails +// at most the first request; a client that does not fails the second too, and every request after +// it, for the life of the connection. +struct SecondRequestOutcome { + bool first_threw{false}; + int first_code{0}; + std::string first_message; + bool second_succeeded{false}; + std::string second_failure; + std::vector reported; ///< What ClientOptions::on_protocol_error saw. +}; + +SecondRequestOutcome send_two_requests( + boost::asio::io_context& io_ctx, + const std::function(const std::string& id)>& answer_first) { + auto transport = std::make_shared(io_ctx.get_executor()); + auto* raw_transport = transport.get(); + + SecondRequestOutcome outcome; + + int answered = 0; + raw_transport->set_on_write([&](std::string_view msg) { + const auto json_msg = nlohmann::json::parse(msg); + if (!json_msg.contains("id")) { + return; + } + const auto id = json_msg.at("id").get(); + if (json_msg.value("method", "") == "initialize") { + raw_transport->enqueue_message(make_result_response(id, make_initialize_result()).dump()); + return; + } + ++answered; + if (answered == 1) { + for (const auto& message : answer_first(id)) { + raw_transport->enqueue_message(message.dump()); + } + } else { + raw_transport->enqueue_message(make_result_response(id, nlohmann::json::object()).dump()); + } + }); + + mcp::ClientOptions options; + // Short enough that a response the client never delivers surfaces as a test failure instead of + // a stalled run. + options.request_timeout = std::chrono::seconds(5); + options.on_protocol_error = [&](const mcp::Error& error) { outcome.reported.push_back(error); }; + mcp::Client client(transport, io_ctx.get_executor(), options); + + boost::asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + static_cast(co_await client.connect("test-client", "0.1")); + + try { + co_await client.ping(); + } catch (const mcp::McpError& error) { + outcome.first_threw = true; + outcome.first_code = error.code(); + outcome.first_message = error.what(); + } + + try { + co_await client.ping(); + outcome.second_succeeded = true; + } catch (const mcp::McpError& error) { + outcome.second_failure = error.what(); + } + + raw_transport->close(); + }, + boost::asio::detached); + + io_ctx.run(); + return outcome; +} + +} // namespace + +// JSON-RPC 2.0 requires `error.message`, but omitting it is an ordinary server-side slip, not an +// attack. It must cost the peer that one response, not the session. +TEST_F(ClientCoreTest, ErrorResponseWithoutMessageDoesNotStopReadLoop) { + const auto outcome = send_two_requests(io_ctx_, [](const std::string& id) { + return std::vector{ + {{"jsonrpc", "2.0"}, {"id", id}, {"error", {{"code", mcp::g_METHOD_NOT_FOUND}}}}}; + }); + + EXPECT_TRUE(outcome.second_succeeded) + << "the session died on one malformed response: " << outcome.second_failure; + ASSERT_TRUE(outcome.first_threw); + EXPECT_EQ(outcome.first_code, mcp::g_METHOD_NOT_FOUND) + << "the error code did not survive the missing message: " << outcome.first_message; +} + +// A peer that serializes absent optionals as explicit null -- the default for Go's encoding/json +// without omitempty -- sends `message: null` rather than omitting the key. +TEST_F(ClientCoreTest, ErrorResponseWithNullMessageDoesNotStopReadLoop) { + const auto outcome = send_two_requests(io_ctx_, [](const std::string& id) { + return std::vector{ + {{"jsonrpc", "2.0"}, + {"id", id}, + {"error", {{"code", mcp::g_METHOD_NOT_FOUND}, {"message", nullptr}, {"data", nullptr}}}}}; + }); + + EXPECT_TRUE(outcome.second_succeeded) + << "the session died on one malformed response: " << outcome.second_failure; + ASSERT_TRUE(outcome.first_threw); + EXPECT_EQ(outcome.first_code, mcp::g_METHOD_NOT_FOUND); +} + +// An `error` member that is not an object yields no code, so the request it answers fails -- but +// it still answers that request rather than leaving the caller to wait out its deadline, and the +// session continues. +TEST_F(ClientCoreTest, ErrorMemberThatIsNotAnObjectFailsOnlyItsOwnRequest) { + const auto outcome = send_two_requests(io_ctx_, [](const std::string& id) { + return std::vector{{{"jsonrpc", "2.0"}, {"id", id}, {"error", "oops"}}}; + }); + + EXPECT_TRUE(outcome.second_succeeded) + << "the session died on one malformed response: " << outcome.second_failure; + ASSERT_TRUE(outcome.first_threw); + EXPECT_EQ(outcome.first_code, mcp::g_INVALID_REQUEST) << "actual: " << outcome.first_message; +} + +// A message the client cannot classify at all is dropped. Dropping it silently would leave an +// application unable to tell a misbehaving peer from a quiet one, so it is reported first. +TEST_F(ClientCoreTest, UndecodableMessageIsReportedAndDropped) { + const auto outcome = send_two_requests(io_ctx_, [](const std::string& id) { + return std::vector{ + nlohmann::json::array({"not", "an", "object"}), + {{"jsonrpc", "2.0"}, {"id", id}, {"result", nlohmann::json::object()}}}; + }); + + EXPECT_TRUE(outcome.second_succeeded) + << "the session died on one malformed message: " << outcome.second_failure; + EXPECT_FALSE(outcome.first_threw) << "actual: " << outcome.first_message; + ASSERT_EQ(outcome.reported.size(), 1U) << "the drop was silent"; + EXPECT_EQ(outcome.reported.front().code, mcp::g_PARSE_ERROR); +} + +// A request from the peer whose id is neither string nor integer can never be answered. Dispatching +// it would bury the failure in a detached coroutine; it is rejected where the drop is reported. +TEST_F(ClientCoreTest, PeerRequestWithUnusableIdDoesNotStopReadLoop) { + const auto outcome = send_two_requests(io_ctx_, [](const std::string& id) { + return std::vector{ + {{"jsonrpc", "2.0"}, {"id", nlohmann::json::object()}, {"method", "ping"}}, + {{"jsonrpc", "2.0"}, {"id", id}, {"result", nlohmann::json::object()}}}; + }); + + EXPECT_TRUE(outcome.second_succeeded) + << "the session died on one malformed request: " << outcome.second_failure; + EXPECT_FALSE(outcome.first_threw) << "actual: " << outcome.first_message; + ASSERT_EQ(outcome.reported.size(), 1U) << "the drop was silent"; + EXPECT_EQ(outcome.reported.front().code, mcp::g_PARSE_ERROR); +} + +// `result` and `error` are not symmetric. A null `result` is a legitimate empty result, but a null +// `error` is the absence of an error -- so selecting on presence alone reads a successful response +// as one carrying both members, and rejects it. The request completes, so nothing hangs and nothing +// times out; the caller is simply handed a protocol violation in place of the result that arrived. +// +// The second request is the control: the same response without the null member. +TEST_F(ClientCoreTest, ResultWithNullErrorMemberIsNotAProtocolViolation) { + const auto outcome = send_two_requests(io_ctx_, [](const std::string& id) { + return std::vector{ + {{"jsonrpc", "2.0"}, {"id", id}, {"result", {{"ok", true}}}, {"error", nullptr}}}; + }); + + EXPECT_FALSE(outcome.first_threw) + << "a successful response was rejected over a null error member: " << outcome.first_message; + EXPECT_TRUE(outcome.second_succeeded) + << "the control response failed too, so the null member is not the difference: " + << outcome.second_failure; + EXPECT_TRUE(outcome.reported.empty()) << "a well-formed response was reported as an error"; +} diff --git a/test/client/client_notifications_test.cpp b/test/client/client_notifications_test.cpp index 48d21ae..a63436c 100644 --- a/test/client/client_notifications_test.cpp +++ b/test/client/client_notifications_test.cpp @@ -7,10 +7,14 @@ #include #include #include +#include #include #include +#include +#include #include #include +#include namespace { @@ -259,3 +263,163 @@ TEST_F(ClientNotificationsTest, UnknownRequestReturnsMethodNotFound) { ASSERT_TRUE(response.contains("error")); EXPECT_EQ(response["error"]["code"], mcp::g_METHOD_NOT_FOUND); } + +namespace { + +struct NotificationSessionOutcome { + bool survived{false}; ///< A request issued after the notifications still completed. + std::string failure; + std::vector reported; ///< What ClientOptions::on_protocol_error saw. +}; + +// Delivers `notifications` to a connected client, then issues one request over the same session. +// That request is the probe: it completes only if the notifications left the session intact. A +// notification the client mishandles by ending the read loop fails it instead. +NotificationSessionOutcome deliver_notifications( + boost::asio::io_context& io_ctx, const std::vector& notifications, + const std::function& register_handlers) { + auto transport = std::make_shared(io_ctx.get_executor()); + auto* raw = transport.get(); + + NotificationSessionOutcome outcome; + + mcp::ClientOptions options; + // Short enough that a swallowed response surfaces as a test failure instead of a stalled run. + options.request_timeout = std::chrono::seconds(5); + options.on_protocol_error = [&](const mcp::Error& error) { outcome.reported.push_back(error); }; + mcp::Client client(transport, io_ctx.get_executor(), options); + register_handlers(client); + + bool delivered = false; + raw->set_on_write([&](std::string_view msg) { + const auto json_msg = nlohmann::json::parse(msg); + if (!json_msg.contains("id")) { + return; + } + const auto id = json_msg.at("id").get(); + if (json_msg.value("method", "") == "initialize") { + raw->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", id}, {"result", make_initialize_result()}} + .dump()); + return; + } + if (!delivered) { + delivered = true; + for (const auto& notification : notifications) { + raw->enqueue_message(notification.dump()); + } + } + raw->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", id}, {"result", nlohmann::json::object()}} + .dump()); + }); + + boost::asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + static_cast(co_await client.connect("test-client", "1.0")); + try { + co_await client.ping(); + outcome.survived = true; + } catch (const mcp::McpError& error) { + outcome.failure = error.what(); + } + raw->close(); + }, + boost::asio::detached); + + io_ctx.run(); + return outcome; +} + +} // namespace + +// A cancelled notification names a request id, and the spec constrains that to a string or an +// integer. A peer that sends anything else has said nothing the client can act on -- and nothing +// that should cost the application its connection. +TEST_F(ClientNotificationsTest, MalformedCancelledNotificationDoesNotStopReadLoop) { + const auto outcome = + deliver_notifications(io_ctx_, + {{{"jsonrpc", "2.0"}, + {"method", "notifications/cancelled"}, + {"params", {{"requestId", nlohmann::json::object()}}}}}, + [](mcp::Client&) {}); + + EXPECT_TRUE(outcome.survived) << "the session died on one malformed notification: " + << outcome.failure; + ASSERT_EQ(outcome.reported.size(), 1U) << "the drop was silent"; + EXPECT_EQ(outcome.reported.front().code, mcp::g_PARSE_ERROR); +} + +// `total` and `message` are optional. A peer that serializes absent optionals as explicit null -- +// the default for Go's encoding/json without omitempty -- still means "not provided", and the +// notification must arrive with those fields empty. +TEST_F(ClientNotificationsTest, ProgressNotificationWithNullOptionalFieldsIsDelivered) { + std::optional seen; + + const auto outcome = deliver_notifications( + io_ctx_, + {{{"jsonrpc", "2.0"}, + {"method", "notifications/progress"}, + {"params", + {{"progressToken", "tok-1"}, {"progress", 0.5}, {"total", nullptr}, {"message", nullptr}}}}}, + [&](mcp::Client& client) { + client.on_progress([&](const mcp::ProgressNotificationParams& params) { seen = params; }); + }); + + EXPECT_TRUE(outcome.survived) << "the session died on one progress notification: " + << outcome.failure; + ASSERT_TRUE(seen.has_value()) << "the progress callback never ran"; + EXPECT_DOUBLE_EQ(seen->progress, 0.5); + EXPECT_FALSE(seen->total.has_value()); + EXPECT_FALSE(seen->message.has_value()); + EXPECT_TRUE(outcome.reported.empty()) << "a well-formed notification was reported as an error"; +} + +// Progress params that are genuinely undecodable are reported against their own cause, so an +// application can tell a peer that sent the notification wrong from a bug in its own callback. +TEST_F(ClientNotificationsTest, MalformedProgressNotificationIsReportedNotFatal) { + bool callback_ran = false; + + const auto outcome = deliver_notifications( + io_ctx_, + {{{"jsonrpc", "2.0"}, + {"method", "notifications/progress"}, + {"params", {{"progressToken", "tok-1"}, {"progress", nullptr}}}}}, + [&](mcp::Client& client) { + client.on_progress([&](const mcp::ProgressNotificationParams&) { callback_ran = true; }); + }); + + EXPECT_TRUE(outcome.survived) << "the session died on one progress notification: " + << outcome.failure; + EXPECT_FALSE(callback_ran) << "an undecodable notification reached the application"; + ASSERT_EQ(outcome.reported.size(), 1U) << "the drop was silent"; + EXPECT_NE(outcome.reported.front().message.find("Malformed progress notification params"), + std::string::npos) + << "actual: " << outcome.reported.front().message; +} + +// The notification callback is application code and it runs on the read loop. A bug in it is the +// application's to fix, not grounds for the SDK to tear down the connection -- but it must not +// vanish either, or the application cannot find the bug. +TEST_F(ClientNotificationsTest, ThrowingNotificationCallbackDoesNotStopReadLoop) { + const auto outcome = deliver_notifications( + io_ctx_, + {{{"jsonrpc", "2.0"}, + {"method", "notifications/message"}, + {"params", nlohmann::json::object()}}}, + [](mcp::Client& client) { + client.on_notification("notifications/message", [](const nlohmann::json&) { + throw std::runtime_error("callback blew up"); + }); + }); + + EXPECT_TRUE(outcome.survived) << "the session died on a throwing application callback: " + << outcome.failure; + ASSERT_EQ(outcome.reported.size(), 1U) << "the throw was silent"; + EXPECT_EQ(outcome.reported.front().code, mcp::g_INTERNAL_ERROR); + EXPECT_NE(outcome.reported.front().message.find("notifications/message"), std::string::npos) + << "actual: " << outcome.reported.front().message; + EXPECT_NE(outcome.reported.front().message.find("callback blew up"), std::string::npos) + << "actual: " << outcome.reported.front().message; +} diff --git a/test/core/capabilities_extensions_test.cpp b/test/core/capabilities_extensions_test.cpp new file mode 100644 index 0000000..1d98e4d --- /dev/null +++ b/test/core/capabilities_extensions_test.cpp @@ -0,0 +1,59 @@ +#include "mcp/protocol/capabilities.hpp" + +#include + +#include + +using json = nlohmann::json; + +TEST(CapabilitiesExtensionsTest, ClientCapabilitiesExtensionsRoundTripsWhenPresent) { + mcp::ClientCapabilities caps; + caps.extensions = std::map{ + {"com.example/foo", json{{"version", 1}}}, + {"com.example/bar", json::object()}, + }; + + json json_obj = caps; + ASSERT_TRUE(json_obj.contains("extensions")); + EXPECT_EQ(json_obj["extensions"]["com.example/foo"]["version"], 1); + EXPECT_TRUE(json_obj["extensions"]["com.example/bar"].is_object()); + + auto round_tripped = json_obj.get(); + ASSERT_TRUE(round_tripped.extensions.has_value()); + EXPECT_EQ(round_tripped.extensions->at("com.example/foo")["version"], 1); + EXPECT_TRUE(round_tripped.extensions->count("com.example/bar")); +} + +TEST(CapabilitiesExtensionsTest, ClientCapabilitiesExtensionsAbsentByDefault) { + mcp::ClientCapabilities caps; + + json json_obj = caps; + EXPECT_FALSE(json_obj.contains("extensions")); + + auto round_tripped = json_obj.get(); + EXPECT_FALSE(round_tripped.extensions.has_value()); +} + +TEST(CapabilitiesExtensionsTest, ServerCapabilitiesExtensionsRoundTripsWhenPresent) { + mcp::ServerCapabilities caps; + caps.extensions = + std::map{{"com.example/baz", json{{"enabled", true}}}}; + + json json_obj = caps; + ASSERT_TRUE(json_obj.contains("extensions")); + EXPECT_EQ(json_obj["extensions"]["com.example/baz"]["enabled"], true); + + auto round_tripped = json_obj.get(); + ASSERT_TRUE(round_tripped.extensions.has_value()); + EXPECT_EQ(round_tripped.extensions->at("com.example/baz")["enabled"], true); +} + +TEST(CapabilitiesExtensionsTest, ServerCapabilitiesExtensionsAbsentByDefault) { + mcp::ServerCapabilities caps; + + json json_obj = caps; + EXPECT_FALSE(json_obj.contains("extensions")); + + auto round_tripped = json_obj.get(); + EXPECT_FALSE(round_tripped.extensions.has_value()); +} diff --git a/test/core/json_matrix_manifest.json b/test/core/json_matrix_manifest.json new file mode 100644 index 0000000..fc28c2e --- /dev/null +++ b/test/core/json_matrix_manifest.json @@ -0,0 +1,146 @@ +{ + "case_count": 1268, + "excluded": { + "ClientCapabilities::TaskRequestsCapability": "no census row governs any of its fields", + "CompleteParams": "no baseline document could be synthesised", + "CompleteReference": "std::variant; from_json dispatches on the 'type' discriminator", + "ContentBlock": "std::variant; from_json dispatches on the 'type' discriminator", + "CreateMessageResult": "no baseline document could be synthesised", + "ElicitRequest": "no baseline document could be synthesised", + "ElicitRequestParams": "std::variant; from_json dispatches on mode", + "EmbeddedResource": "no baseline document could be synthesised", + "InitializeRequest": "no baseline document could be synthesised", + "InitializeResult": "no baseline document could be synthesised", + "JSONRPCMessage": "std::variant; from_json dispatches on which key is present", + "JSONRPCResponse": "std::variant; from_json dispatches on which key is present", + "PrimitiveSchemaDefinition": "base of EnumSchema; exercised through its derived type", + "PromptMessage": "no baseline document could be synthesised", + "RequestId": "std::variant of string and integer; it has no fields to vary", + "ResourceContents": "std::variant; from_json dispatches on text vs blob", + "SamplingMessage": "no baseline document could be synthesised", + "SamplingMessageContent": "std::variant; from_json accepts a block or a list of blocks", + "ServerCapabilities::TaskRequestsCapability": "no census row governs any of its fields" + }, + "field_count": 317, + "gtest_count": 118, + "mode_count": 4, + "types": [ + "Annotations", + "AudioContent", + "BlobResourceContents", + "CallToolParams", + "CallToolRequest", + "CallToolResult", + "CancelTaskRequest", + "CancelTaskRequestParams", + "CancelTaskResult", + "CancelledNotification", + "CancelledNotificationParams", + "ClientCapabilities", + "ClientCapabilities::ElicitationCapability", + "ClientCapabilities::RootsCapability", + "ClientCapabilities::SamplingCapability", + "ClientCapabilities::TaskRequestsCapability::ElicitationTaskCapability", + "ClientCapabilities::TaskRequestsCapability::SamplingTaskCapability", + "ClientCapabilities::TasksCapability", + "CompleteContext", + "CompleteParamsArgument", + "CompleteResult", + "CompletionResultDetails", + "CreateMessageRequest", + "CreateMessageRequestParams", + "CreateTaskResult", + "DiscoverRequest", + "DiscoverResult", + "ElicitRequestFormParams", + "ElicitRequestURLParams", + "ElicitResult", + "ElicitationCompleteNotification", + "ElicitationCompleteNotificationParams", + "EnumSchema", + "Error", + "GetPromptRequest", + "GetPromptRequestParams", + "GetPromptResult", + "GetTaskPayloadRequest", + "GetTaskPayloadRequestParams", + "GetTaskPayloadResult", + "GetTaskRequest", + "GetTaskRequestParams", + "GetTaskResult", + "Icon", + "ImageContent", + "Implementation", + "InitializedNotification", + "JSONRPCErrorResponse", + "JSONRPCNotification", + "JSONRPCRequest", + "JSONRPCResultResponse", + "ListPromptsRequest", + "ListPromptsResult", + "ListResourceTemplatesRequest", + "ListResourceTemplatesRequestParams", + "ListResourceTemplatesResult", + "ListResourcesRequest", + "ListResourcesRequestParams", + "ListResourcesResult", + "ListRootsRequest", + "ListRootsResult", + "ListTasksRequest", + "ListTasksRequestParams", + "ListTasksResult", + "ListToolsRequest", + "ListToolsRequestParams", + "ListToolsResult", + "LoggingMessageNotification", + "LoggingMessageNotificationParams", + "ModelHint", + "ModelPreferences", + "PaginatedRequestParams", + "PingRequest", + "ProgressNotification", + "ProgressNotificationParams", + "Prompt", + "PromptArgument", + "PromptListChangedNotification", + "PromptReference", + "ReadResourceRequest", + "ReadResourceRequestParams", + "ReadResourceResult", + "RelatedTaskMetadata", + "Resource", + "ResourceLink", + "ResourceListChangedNotification", + "ResourceSubscribeParams", + "ResourceTemplate", + "ResourceTemplateReference", + "ResourceUnsubscribeParams", + "ResourceUpdatedNotification", + "ResourceUpdatedNotificationParams", + "Root", + "RootsListChangedNotification", + "ServerCapabilities", + "ServerCapabilities::PromptsCapability", + "ServerCapabilities::ResourcesCapability", + "ServerCapabilities::TaskRequestsCapability::ToolsTaskCapability", + "ServerCapabilities::TasksCapability", + "ServerCapabilities::ToolsCapability", + "SetLevelRequest", + "SetLevelRequestParams", + "SubscribeRequest", + "TaskData", + "TaskMetadata", + "TaskStatusNotification", + "TaskStatusNotificationParams", + "TextContent", + "TextResourceContents", + "Tool", + "ToolAnnotations", + "ToolChoice", + "ToolExecution", + "ToolListChangedNotification", + "ToolResultContent", + "ToolUseContent", + "UnsubscribeRequest" + ] +} diff --git a/test/core/json_peer_input_matrix_test.cpp b/test/core/json_peer_input_matrix_test.cpp new file mode 100644 index 0000000..1f09d5c --- /dev/null +++ b/test/core/json_peer_input_matrix_test.cpp @@ -0,0 +1,1846 @@ +// Generated by scripts/gen_json_matrix.py from scripts/json_census.py. +// Do not edit by hand: regenerate with `python3 scripts/gen_json_matrix.py`. +// +// Every peer-deserialisable protocol type, every field, four modes: the key +// absent, the key present and explicitly null, the key holding the wrong JSON +// type, and the key holding an oversized value of the right type. +// +// Each case asserts what the census predicts, so a disagreement between the +// code and the census fails here rather than in a user's session. Serialising +// an absent optional as an explicit null is the default for Go's +// encoding/json without omitempty and for a naively dumped Python dataclass, +// which puts the null column within reach of a careless peer rather than only +// a hostile one. +// +// scripts/check_json_matrix.py runs as a build step. It fails when a protocol +// type has a from_json but is neither in this matrix nor excluded on the +// record, so the matrix cannot quietly fall behind the protocol, and when this +// file differs from what the generator produces from the current sources. +// +// WHAT THIS CORPUS IS, AND WHAT IT IS NOT +// +// It is a consistency oracle, not a correctness one. Every expectation is +// derived by static analysis of the guard construct at the decode site, and +// the test then runs the decoder. A case fails when runtime and static +// analysis disagree. It cannot fail because a field is wrong per the MCP +// spec. Nothing here is produced by executing the SDK, so the two sides are +// independent -- but agreement means consistency, not conformance. +// +// It follows that once a decode site is fixed and this file is regenerated, +// the matching cases pass BY CONSTRUCTION: the census re-reads the fixed +// source and predicts the new behaviour. Green here is not evidence that a +// field decodes correctly. That evidence lives in the hand-written, +// red-first tests in test/core/protocol_test.cpp and +// test/server/server_handlers_test.cpp. +// +// A TRUE null_throws ON AN OPTIONAL MEMBER RECORDS A DEFECT, NOT A SPEC +// +// The tuple (absent=false, null=true, wrong=true, oversized=false) describes +// a member that tolerates being absent but throws when a peer sends it as an +// explicit null. That is the defect class this project has repeatedly been +// caught by, and scripts/json_census_dispositions.json dispositions it "fix". +// 66 of the 317 field entries below still carry it; there were 87 before the +// explicit-null decoding fixes. Those rows pin behaviour as it is so the +// suite stays green. They do not endorse it. Fixing one makes the build check +// report this file stale until it is regenerated with the fix. The check reads +// the other direction as a REGRESSION: a row whose decoder now throws on an +// explicit null or on an absent key, where this file says it tolerates it. +// The generator refuses to write such a row unless it is named with +// --accept-regression. Never relax a fix to satisfy this file. +// +// This column is also blind to the other half of that class. A member that +// decodes an explicit null into an ENGAGED optional throws nothing, so it is +// recorded as null_throws = false whether or not it then re-encodes a member +// the peer never sent. _meta and annotations were corrupt in exactly that way +// and not one case below changed when they were fixed. Seeing that class +// needs a round-trip check, which this corpus does not have. + +#include + +#include + +#include + +#include + +namespace { + +using nlohmann::json; + +// How to pick a value a field does not accept. A variant of string and integer +// needs an object: "the other scalar type" is still something it accepts. +enum WrongHint { + kWrongFromValue = 0, + kWrongForceObject = 1 +}; + +json wrong_typed(const json& value, int hint) { + if (hint == kWrongForceObject) { + return json{{"neither", "a string nor an integer"}}; + } + if (value.is_boolean()) { + return "not-a-bool"; + } + if (value.is_number()) { + return "not-a-number"; + } + if (value.is_string()) { + return 12345; + } + if (value.is_array()) { + return json{{"not", "an array"}}; + } + return json::array({"not an object"}); +} + +// A very large value that still has the type the field accepts. An oversized +// array stays an array *of its own element type*: filling it with strings +// would make it a wrong-type case wearing an oversized label, and the throw +// would then be read as a size limit that does not exist. +json oversized(const json& value) { + if (value.is_boolean()) { + return value; + } + if (value.is_number()) { + return json(1000000000000000000LL); + } + if (value.is_string()) { + return std::string(65536, 'A'); + } + if (value.is_array()) { + if (value.empty()) { + return value; + } + json grown = json::array(); + for (int i = 0; i < 512; ++i) { + grown.push_back(value.front()); + } + return grown; + } + // An object grows by keys the serialiser ignores, so its shape is intact. + json grown = value.is_object() ? value : json::object(); + for (int i = 0; i < 256; ++i) { + grown["_pad" + std::to_string(i)] = std::string(256, 'A'); + } + return grown; +} + +// Parses `doc` into T and reports whether that threw. +template +bool throws_on(const json& doc) { + try { + static_cast(doc.get()); + return false; + } catch (const std::exception&) { + return true; + } +} + +struct FieldExpect { + const char* key; + const char* present; // a value the field accepts, as JSON text + int wrong_hint; + bool absent_throws; + bool null_throws; + bool wrong_throws; + bool oversized_throws; +}; + +// Applies all four modes to one field and checks each against the census. +template +void check_field(const json& baseline, const FieldExpect& f) { + const json present = json::parse(f.present); + + json absent = baseline; + absent.erase(f.key); + { + SCOPED_TRACE("absent"); + EXPECT_EQ(throws_on(absent), f.absent_throws); + } + json with_null = baseline; + with_null[f.key] = nullptr; + { + SCOPED_TRACE("null"); + EXPECT_EQ(throws_on(with_null), f.null_throws); + } + json with_wrong = baseline; + with_wrong[f.key] = wrong_typed(present, f.wrong_hint); + { + SCOPED_TRACE("wrong_type"); + EXPECT_EQ(throws_on(with_wrong), f.wrong_throws); + } + json with_big = baseline; + with_big[f.key] = oversized(present); + { + SCOPED_TRACE("oversized"); + EXPECT_EQ(throws_on(with_big), f.oversized_throws); + } +} + +} // namespace + +TEST(JsonPeerInputMatrix, Annotations) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"audience", R"json(["user"])json", 0, false, true, true, false}, + {"lastModified", R"json("x")json", 0, false, true, true, false}, + {"priority", R"json(1.0)json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Annotations.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, AudioContent) { + const json baseline = json::parse(R"json({"type":"audio","data":"x","mimeType":"x"})json"); + static const FieldExpect fields[] = { + {"data", R"json("x")json", 0, true, true, true, false}, + {"mimeType", R"json("x")json", 0, true, true, true, false}, + {"type", R"json("audio")json", 0, true, true, true, true}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("AudioContent.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, BlobResourceContents) { + const json baseline = json::parse(R"json({"uri":"x","blob":"x"})json"); + static const FieldExpect fields[] = { + {"blob", R"json("x")json", 0, true, true, true, false}, + {"mimeType", R"json("x")json", 0, false, true, true, false}, + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("BlobResourceContents.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CallToolParams) { + const json baseline = json::parse(R"json({"name":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CallToolParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CallToolRequest) { + const json baseline = json::parse(R"json({"method":"tools/call","params":{"name":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("tools/call")json", 0, true, true, true, false}, + {"params", R"json({"name":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CallToolRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CallToolResult) { + const json baseline = json::parse(R"json({"content":[]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"content", R"json([])json", 0, true, true, true, false}, + {"isError", R"json(true)json", 0, false, false, true, false}, + {"structuredContent", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CallToolResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CancelTaskRequest) { + const json baseline = json::parse(R"json({"method":"tasks/cancel","params":{"id":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("tasks/cancel")json", 0, true, true, true, false}, + {"params", R"json({"id":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CancelTaskRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CancelTaskRequestParams) { + const json baseline = json::parse(R"json({"id":"x"})json"); + static const FieldExpect fields[] = { + {"id", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CancelTaskRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CancelTaskResult) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CancelTaskResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CancelledNotification) { + const json baseline = + json::parse(R"json({"method":"notifications/cancelled","params":{"requestId":1}})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/cancelled")json", 0, true, true, true, false}, + {"params", R"json({"requestId":1})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CancelledNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CancelledNotificationParams) { + const json baseline = json::parse(R"json({"requestId":1})json"); + static const FieldExpect fields[] = { + {"reason", R"json("x")json", 0, false, false, true, false}, + {"requestId", R"json(1)json", 1, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CancelledNotificationParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ClientCapabilities) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"experimental", R"json({})json", 0, false, false, false, false}, + {"extensions", R"json({})json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ClientCapabilities.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ClientCapabilities_ElicitationCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"form", R"json({})json", 0, false, false, false, false}, + {"url", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ClientCapabilities::ElicitationCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ClientCapabilities_RootsCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"listChanged", R"json(true)json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ClientCapabilities::RootsCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ClientCapabilities_SamplingCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"context", R"json({})json", 0, false, false, false, false}, + {"tools", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ClientCapabilities::SamplingCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ClientCapabilities_TaskRequestsCapability_ElicitationTaskCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"create", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE( + std::string("ClientCapabilities::TaskRequestsCapability::ElicitationTaskCapability.") + + f.key); + check_field( + baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ClientCapabilities_TaskRequestsCapability_SamplingTaskCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"createMessage", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE( + std::string("ClientCapabilities::TaskRequestsCapability::SamplingTaskCapability.") + f.key); + check_field(baseline, + f); + } +} + +TEST(JsonPeerInputMatrix, ClientCapabilities_TasksCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"cancel", R"json({})json", 0, false, false, false, false}, + {"list", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ClientCapabilities::TasksCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CompleteContext) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"arguments", R"json({})json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CompleteContext.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CompleteParamsArgument) { + const json baseline = json::parse(R"json({"name":"x","value":"x"})json"); + static const FieldExpect fields[] = { + {"name", R"json("x")json", 0, true, true, true, false}, + {"value", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CompleteParamsArgument.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CompleteResult) { + const json baseline = json::parse(R"json({"completion":{"values":["x"]}})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"completion", R"json({"values":["x"]})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CompleteResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CompletionResultDetails) { + const json baseline = json::parse(R"json({"values":["x"]})json"); + static const FieldExpect fields[] = { + {"hasMore", R"json(true)json", 0, false, true, true, false}, + {"total", R"json(1)json", 0, false, true, true, false}, + {"values", R"json(["x"])json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CompletionResultDetails.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CreateMessageRequest) { + const json baseline = json::parse( + R"json({"method":"sampling/createMessage","params":{"messages":[],"maxTokens":1}})json"); + static const FieldExpect fields[] = { + {"method", R"json("sampling/createMessage")json", 0, true, true, true, false}, + {"params", R"json({"messages":[],"maxTokens":1})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CreateMessageRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CreateMessageRequestParams) { + const json baseline = json::parse(R"json({"messages":[],"maxTokens":1})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"includeContext", R"json("x")json", 0, false, true, true, false}, + {"maxTokens", R"json(1)json", 0, true, true, true, false}, + {"messages", R"json([])json", 0, true, true, true, false}, + {"metadata", R"json({})json", 0, false, false, false, false}, + {"modelPreferences", R"json({})json", 0, false, false, false, false}, + {"stopSequences", R"json(["x"])json", 0, false, true, true, false}, + {"systemPrompt", R"json("x")json", 0, false, true, true, false}, + {"temperature", R"json(1.0)json", 0, false, true, true, false}, + {"toolChoice", R"json({})json", 0, false, false, false, false}, + {"tools", R"json([{"name":"x","inputSchema":{}}])json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CreateMessageRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, CreateTaskResult) { + const json baseline = json::parse( + R"json({"task":{"taskId":"x","status":"cancelled","createdAt":"x","lastUpdatedAt":"x","ttl":1}})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"task", + R"json({"taskId":"x","status":"cancelled","createdAt":"x","lastUpdatedAt":"x","ttl":1})json", + 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("CreateTaskResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, DiscoverRequest) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("DiscoverRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, DiscoverResult) { + const json baseline = + json::parse(R"json({"resultType":"complete","supportedVersions":["x"],"capabilities":{}})json"); + static const FieldExpect fields[] = { + {"cacheScope", R"json("public")json", 0, false, false, false, false}, + {"capabilities", R"json({})json", 0, true, false, false, false}, + {"instructions", R"json("x")json", 0, false, true, true, false}, + {"resultType", R"json("complete")json", 0, true, true, true, false}, + {"supportedVersions", R"json(["x"])json", 0, true, true, true, false}, + {"ttlMs", R"json(1)json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("DiscoverResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ElicitRequestFormParams) { + const json baseline = json::parse(R"json({"mode":"form","message":"x","requestedSchema":{}})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"message", R"json("x")json", 0, true, true, true, false}, + {"mode", R"json("form")json", 0, true, true, true, false}, + {"requestedSchema", R"json({})json", 0, true, false, false, false}, + {"task", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ElicitRequestFormParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ElicitRequestURLParams) { + const json baseline = + json::parse(R"json({"mode":"url","message":"x","url":"x","elicitationId":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"elicitationId", R"json("x")json", 0, true, true, true, false}, + {"message", R"json("x")json", 0, true, true, true, false}, + {"mode", R"json("url")json", 0, true, true, true, false}, + {"task", R"json({})json", 0, false, false, false, false}, + {"url", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ElicitRequestURLParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ElicitResult) { + const json baseline = json::parse(R"json({"action":"accept"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"action", R"json("accept")json", 0, true, false, false, false}, + {"content", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ElicitResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ElicitationCompleteNotification) { + const json baseline = json::parse( + R"json({"method":"notifications/elicitation/complete","params":{"requestId":1}})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/elicitation/complete")json", 0, true, true, true, false}, + {"params", R"json({"requestId":1})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ElicitationCompleteNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ElicitationCompleteNotificationParams) { + const json baseline = json::parse(R"json({"requestId":1})json"); + static const FieldExpect fields[] = { + {"requestId", R"json(1)json", 1, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ElicitationCompleteNotificationParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, EnumSchema) { + const json baseline = json::parse(R"json({"type":"x"})json"); + static const FieldExpect fields[] = { + {"description", R"json("x")json", 0, false, true, true, false}, + {"enum", R"json(["x"])json", 0, false, true, true, false}, + {"title", R"json("x")json", 0, false, true, true, false}, + {"type", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("EnumSchema.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, Error) { + const json baseline = json::parse(R"json({"code":1})json"); + static const FieldExpect fields[] = { + {"code", R"json(1)json", 0, true, true, true, false}, + {"data", R"json({})json", 0, false, false, false, false}, + {"message", R"json("x")json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Error.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetPromptRequest) { + const json baseline = json::parse(R"json({"method":"prompts/get","params":{"name":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("prompts/get")json", 0, true, true, true, false}, + {"params", R"json({"name":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetPromptRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetPromptRequestParams) { + const json baseline = json::parse(R"json({"name":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"arguments", R"json({})json", 0, false, false, true, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetPromptRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetPromptResult) { + const json baseline = json::parse(R"json({"messages":[]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"description", R"json("x")json", 0, false, true, true, false}, + {"messages", R"json([])json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetPromptResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetTaskPayloadRequest) { + const json baseline = json::parse(R"json({"method":"tasks/getPayload","params":{"id":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("tasks/getPayload")json", 0, true, true, true, false}, + {"params", R"json({"id":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetTaskPayloadRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetTaskPayloadRequestParams) { + const json baseline = json::parse(R"json({"id":"x"})json"); + static const FieldExpect fields[] = { + {"id", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetTaskPayloadRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetTaskPayloadResult) { + const json baseline = json::parse(R"json({"payload":{}})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"payload", R"json({})json", 0, true, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetTaskPayloadResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetTaskRequest) { + const json baseline = json::parse(R"json({"method":"tasks/get","params":{"id":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("tasks/get")json", 0, true, true, true, false}, + {"params", R"json({"id":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetTaskRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetTaskRequestParams) { + const json baseline = json::parse(R"json({"id":"x"})json"); + static const FieldExpect fields[] = { + {"id", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetTaskRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, GetTaskResult) { + const json baseline = json::parse( + R"json({"taskId":"x","status":"cancelled","createdAt":"x","lastUpdatedAt":"x","ttl":1})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"createdAt", R"json("x")json", 0, true, true, true, false}, + {"lastUpdatedAt", R"json("x")json", 0, true, true, true, false}, + {"pollInterval", R"json(1)json", 0, false, true, true, false}, + {"status", R"json("cancelled")json", 0, true, false, false, false}, + {"statusMessage", R"json("x")json", 0, false, true, true, false}, + {"taskId", R"json("x")json", 0, true, true, true, false}, + {"ttl", R"json(1)json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("GetTaskResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, Icon) { + const json baseline = json::parse(R"json({"src":"x"})json"); + static const FieldExpect fields[] = { + {"mimeType", R"json("x")json", 0, false, true, true, false}, + {"sizes", R"json(["x"])json", 0, false, true, true, false}, + {"src", R"json("x")json", 0, true, true, true, false}, + {"theme", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Icon.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ImageContent) { + const json baseline = json::parse(R"json({"type":"image","data":"x","mimeType":"x"})json"); + static const FieldExpect fields[] = { + {"data", R"json("x")json", 0, true, true, true, false}, + {"mimeType", R"json("x")json", 0, true, true, true, false}, + {"type", R"json("image")json", 0, true, true, true, true}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ImageContent.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, Implementation) { + const json baseline = json::parse(R"json({"name":"x","version":"x"})json"); + static const FieldExpect fields[] = { + {"description", R"json("x")json", 0, false, false, true, false}, + {"icons", R"json([{"src":"x"}])json", 0, false, false, true, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"title", R"json("x")json", 0, false, false, true, false}, + {"version", R"json("x")json", 0, true, true, true, false}, + {"websiteUrl", R"json("x")json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Implementation.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, InitializedNotification) { + const json baseline = json::parse(R"json({"method":"notifications/initialized"})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/initialized")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("InitializedNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, JSONRPCErrorResponse) { + const json baseline = json::parse(R"json({"error":{"code":1},"jsonrpc":"2.0"})json"); + static const FieldExpect fields[] = { + {"error", R"json({"code":1})json", 0, true, true, true, false}, + {"id", R"json(1)json", 1, false, false, true, false}, + {"jsonrpc", R"json("2.0")json", 0, true, true, true, true}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("JSONRPCErrorResponse.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, JSONRPCNotification) { + const json baseline = json::parse(R"json({"jsonrpc":"2.0","method":"x"})json"); + static const FieldExpect fields[] = { + {"jsonrpc", R"json("2.0")json", 0, true, true, true, true}, + {"method", R"json("x")json", 0, true, true, true, false}, + {"params", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("JSONRPCNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, JSONRPCRequest) { + const json baseline = json::parse(R"json({"id":1,"jsonrpc":"2.0","method":"x"})json"); + static const FieldExpect fields[] = { + {"id", R"json(1)json", 1, true, true, true, false}, + {"jsonrpc", R"json("2.0")json", 0, true, true, true, true}, + {"method", R"json("x")json", 0, true, true, true, false}, + {"params", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("JSONRPCRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, JSONRPCResultResponse) { + const json baseline = json::parse(R"json({"id":1,"jsonrpc":"2.0","result":{}})json"); + static const FieldExpect fields[] = { + {"id", R"json(1)json", 1, true, true, true, false}, + {"jsonrpc", R"json("2.0")json", 0, true, true, true, true}, + {"result", R"json({})json", 0, true, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("JSONRPCResultResponse.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListPromptsRequest) { + const json baseline = json::parse(R"json({"method":"prompts/list"})json"); + static const FieldExpect fields[] = { + {"method", R"json("prompts/list")json", 0, true, true, true, false}, + {"params", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListPromptsRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListPromptsResult) { + const json baseline = json::parse(R"json({"prompts":[{"name":"x"}]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"nextCursor", R"json("x")json", 0, false, true, true, false}, + {"prompts", R"json([{"name":"x"}])json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListPromptsResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListResourceTemplatesRequest) { + const json baseline = json::parse(R"json({"method":"resources/templates/list"})json"); + static const FieldExpect fields[] = { + {"method", R"json("resources/templates/list")json", 0, true, true, true, false}, + {"params", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListResourceTemplatesRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListResourceTemplatesRequestParams) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"cursor", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListResourceTemplatesRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListResourceTemplatesResult) { + const json baseline = + json::parse(R"json({"resourceTemplates":[{"uriTemplate":"x","name":"x"}]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"nextCursor", R"json("x")json", 0, false, true, true, false}, + {"resourceTemplates", R"json([{"uriTemplate":"x","name":"x"}])json", 0, true, true, true, + false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListResourceTemplatesResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListResourcesRequest) { + const json baseline = json::parse(R"json({"method":"resources/list"})json"); + static const FieldExpect fields[] = { + {"method", R"json("resources/list")json", 0, true, true, true, false}, + {"params", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListResourcesRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListResourcesRequestParams) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"cursor", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListResourcesRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListResourcesResult) { + const json baseline = json::parse(R"json({"resources":[{"uri":"x","name":"x"}]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"nextCursor", R"json("x")json", 0, false, true, true, false}, + {"resources", R"json([{"uri":"x","name":"x"}])json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListResourcesResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListRootsRequest) { + const json baseline = json::parse(R"json({"method":"roots/list"})json"); + static const FieldExpect fields[] = { + {"method", R"json("roots/list")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListRootsRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListRootsResult) { + const json baseline = json::parse(R"json({"roots":[{"uri":"x"}]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"roots", R"json([{"uri":"x"}])json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListRootsResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListTasksRequest) { + const json baseline = json::parse(R"json({"method":"tasks/list"})json"); + static const FieldExpect fields[] = { + {"method", R"json("tasks/list")json", 0, true, true, true, false}, + {"params", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListTasksRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListTasksRequestParams) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"cursor", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListTasksRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListTasksResult) { + const json baseline = json::parse( + R"json({"tasks":[{"taskId":"x","status":"cancelled","createdAt":"x","lastUpdatedAt":"x","ttl":1}]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"nextCursor", R"json("x")json", 0, false, true, true, false}, + {"tasks", + R"json([{"taskId":"x","status":"cancelled","createdAt":"x","lastUpdatedAt":"x","ttl":1}])json", + 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListTasksResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListToolsRequest) { + const json baseline = json::parse(R"json({"method":"tools/list"})json"); + static const FieldExpect fields[] = { + {"method", R"json("tools/list")json", 0, true, true, true, false}, + {"params", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListToolsRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListToolsRequestParams) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"cursor", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListToolsRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ListToolsResult) { + const json baseline = json::parse(R"json({"tools":[{"name":"x","inputSchema":{}}]})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"nextCursor", R"json("x")json", 0, false, true, true, false}, + {"tools", R"json([{"name":"x","inputSchema":{}}])json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ListToolsResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, LoggingMessageNotification) { + const json baseline = json::parse( + R"json({"method":"notifications/message","params":{"level":"emergency","data":{}}})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/message")json", 0, true, true, true, false}, + {"params", R"json({"level":"emergency","data":{}})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("LoggingMessageNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, LoggingMessageNotificationParams) { + const json baseline = json::parse(R"json({"level":"emergency","data":{}})json"); + static const FieldExpect fields[] = { + {"data", R"json({})json", 0, true, false, false, false}, + {"level", R"json("emergency")json", 0, true, false, false, false}, + {"logger", R"json("x")json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("LoggingMessageNotificationParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ModelHint) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"name", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ModelHint.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ModelPreferences) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"costPriority", R"json(1.0)json", 0, false, true, true, false}, + {"hints", R"json([{}])json", 0, false, true, true, false}, + {"intelligencePriority", R"json(1.0)json", 0, false, true, true, false}, + {"speedPriority", R"json(1.0)json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ModelPreferences.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, PaginatedRequestParams) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"cursor", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("PaginatedRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, PingRequest) { + const json baseline = json::parse(R"json({"method":"ping"})json"); + static const FieldExpect fields[] = { + {"method", R"json("ping")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("PingRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ProgressNotification) { + const json baseline = json::parse( + R"json({"method":"notifications/progress","params":{"progressToken":1,"progress":1.0}})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/progress")json", 0, true, true, true, false}, + {"params", R"json({"progressToken":1,"progress":1.0})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ProgressNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ProgressNotificationParams) { + const json baseline = json::parse(R"json({"progressToken":1,"progress":1.0})json"); + static const FieldExpect fields[] = { + {"message", R"json("x")json", 0, false, false, true, false}, + {"progress", R"json(1.0)json", 0, true, true, true, false}, + {"progressToken", R"json(1)json", 1, true, true, true, false}, + {"total", R"json(1.0)json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ProgressNotificationParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, Prompt) { + const json baseline = json::parse(R"json({"name":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"arguments", R"json([{"name":"x"}])json", 0, false, true, true, false}, + {"description", R"json("x")json", 0, false, true, true, false}, + {"icons", R"json([{"src":"x"}])json", 0, false, true, true, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"title", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Prompt.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, PromptArgument) { + const json baseline = json::parse(R"json({"name":"x"})json"); + static const FieldExpect fields[] = { + {"description", R"json("x")json", 0, false, true, true, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"required", R"json(true)json", 0, false, true, true, false}, + {"title", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("PromptArgument.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, PromptListChangedNotification) { + const json baseline = json::parse(R"json({"method":"notifications/prompts/list_changed"})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/prompts/list_changed")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("PromptListChangedNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, PromptReference) { + const json baseline = json::parse(R"json({"name":"x","type":"ref/prompt"})json"); + static const FieldExpect fields[] = { + {"name", R"json("x")json", 0, true, true, true, false}, + {"type", R"json("ref/prompt")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("PromptReference.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ReadResourceRequest) { + const json baseline = json::parse(R"json({"method":"resources/read","params":{"uri":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("resources/read")json", 0, true, true, true, false}, + {"params", R"json({"uri":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ReadResourceRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ReadResourceRequestParams) { + const json baseline = json::parse(R"json({"uri":"x"})json"); + static const FieldExpect fields[] = { + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ReadResourceRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ReadResourceResult) { + const json baseline = json::parse(R"json({"contents":[]})json"); + static const FieldExpect fields[] = { + {"contents", R"json([])json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ReadResourceResult.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, RelatedTaskMetadata) { + const json baseline = json::parse(R"json({"id":"x"})json"); + static const FieldExpect fields[] = { + {"id", R"json("x")json", 0, true, true, true, false}, + {"title", R"json("x")json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("RelatedTaskMetadata.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, Resource) { + const json baseline = json::parse(R"json({"uri":"x","name":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"annotations", R"json({})json", 0, false, false, false, false}, + {"description", R"json("x")json", 0, false, false, true, false}, + {"icons", R"json([{"src":"x"}])json", 0, false, false, true, false}, + {"mimeType", R"json("x")json", 0, false, false, true, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"size", R"json(1)json", 0, false, false, true, false}, + {"title", R"json("x")json", 0, false, false, true, false}, + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Resource.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceLink) { + const json baseline = json::parse(R"json({"type":"resource_link","uri":"x","name":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"annotations", R"json({})json", 0, false, false, false, false}, + {"description", R"json("x")json", 0, false, true, true, false}, + {"icons", R"json([{"src":"x"}])json", 0, false, true, true, false}, + {"mimeType", R"json("x")json", 0, false, true, true, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"size", R"json(1)json", 0, false, true, true, false}, + {"title", R"json("x")json", 0, false, true, true, false}, + {"type", R"json("resource_link")json", 0, true, true, true, true}, + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceLink.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceListChangedNotification) { + const json baseline = json::parse(R"json({"method":"notifications/resources/list_changed"})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/resources/list_changed")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceListChangedNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceSubscribeParams) { + const json baseline = json::parse(R"json({"uri":"x"})json"); + static const FieldExpect fields[] = { + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceSubscribeParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceTemplate) { + const json baseline = json::parse(R"json({"uriTemplate":"x","name":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"annotations", R"json({})json", 0, false, false, false, false}, + {"description", R"json("x")json", 0, false, false, true, false}, + {"icons", R"json([{"src":"x"}])json", 0, false, false, true, false}, + {"mimeType", R"json("x")json", 0, false, false, true, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"title", R"json("x")json", 0, false, false, true, false}, + {"uriTemplate", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceTemplate.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceTemplateReference) { + const json baseline = json::parse(R"json({"type":"ref/resource","uri":"x"})json"); + static const FieldExpect fields[] = { + {"type", R"json("ref/resource")json", 0, true, true, true, false}, + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceTemplateReference.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceUnsubscribeParams) { + const json baseline = json::parse(R"json({"uri":"x"})json"); + static const FieldExpect fields[] = { + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceUnsubscribeParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceUpdatedNotification) { + const json baseline = + json::parse(R"json({"method":"notifications/resources/updated","params":{"uri":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/resources/updated")json", 0, true, true, true, false}, + {"params", R"json({"uri":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceUpdatedNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ResourceUpdatedNotificationParams) { + const json baseline = json::parse(R"json({"uri":"x"})json"); + static const FieldExpect fields[] = { + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ResourceUpdatedNotificationParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, Root) { + const json baseline = json::parse(R"json({"uri":"x"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"name", R"json("x")json", 0, false, true, true, false}, + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Root.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, RootsListChangedNotification) { + const json baseline = json::parse(R"json({"method":"notifications/roots/list_changed"})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/roots/list_changed")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("RootsListChangedNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ServerCapabilities) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"completions", R"json({})json", 0, false, false, false, false}, + {"experimental", R"json({})json", 0, false, false, false, false}, + {"extensions", R"json({})json", 0, false, true, true, false}, + {"logging", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ServerCapabilities.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ServerCapabilities_PromptsCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"listChanged", R"json(true)json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ServerCapabilities::PromptsCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ServerCapabilities_ResourcesCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"listChanged", R"json(true)json", 0, false, true, true, false}, + {"subscribe", R"json(true)json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ServerCapabilities::ResourcesCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ServerCapabilities_TaskRequestsCapability_ToolsTaskCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"call", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ServerCapabilities::TaskRequestsCapability::ToolsTaskCapability.") + + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ServerCapabilities_TasksCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"cancel", R"json({})json", 0, false, false, false, false}, + {"list", R"json({})json", 0, false, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ServerCapabilities::TasksCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ServerCapabilities_ToolsCapability) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"listChanged", R"json(true)json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ServerCapabilities::ToolsCapability.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, SetLevelRequest) { + const json baseline = + json::parse(R"json({"method":"logging/setLevel","params":{"level":"emergency"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("logging/setLevel")json", 0, true, true, true, false}, + {"params", R"json({"level":"emergency"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("SetLevelRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, SetLevelRequestParams) { + const json baseline = json::parse(R"json({"level":"emergency"})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"level", R"json("emergency")json", 0, true, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("SetLevelRequestParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, SubscribeRequest) { + const json baseline = + json::parse(R"json({"method":"resources/subscribe","params":{"uri":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("resources/subscribe")json", 0, true, true, true, false}, + {"params", R"json({"uri":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("SubscribeRequest.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, TaskData) { + const json baseline = json::parse( + R"json({"taskId":"x","status":"cancelled","createdAt":"x","lastUpdatedAt":"x","ttl":1})json"); + static const FieldExpect fields[] = { + {"createdAt", R"json("x")json", 0, true, true, true, false}, + {"lastUpdatedAt", R"json("x")json", 0, true, true, true, false}, + {"pollInterval", R"json(1)json", 0, false, true, true, false}, + {"status", R"json("cancelled")json", 0, true, false, false, false}, + {"statusMessage", R"json("x")json", 0, false, true, true, false}, + {"taskId", R"json("x")json", 0, true, true, true, false}, + {"ttl", R"json(1)json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("TaskData.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, TaskMetadata) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"relatedTasks", R"json([{"id":"x"}])json", 0, false, false, true, false}, + {"ttl", R"json(1)json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("TaskMetadata.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, TaskStatusNotification) { + const json baseline = json::parse( + R"json({"method":"notifications/tasks/status","params":{"id":"x","status":"cancelled"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/tasks/status")json", 0, true, true, true, false}, + {"params", R"json({"id":"x","status":"cancelled"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("TaskStatusNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, TaskStatusNotificationParams) { + const json baseline = json::parse(R"json({"id":"x","status":"cancelled"})json"); + static const FieldExpect fields[] = { + {"id", R"json("x")json", 0, true, true, true, false}, + {"message", R"json("x")json", 0, false, false, true, false}, + {"metadata", R"json({})json", 0, false, false, false, false}, + {"status", R"json("cancelled")json", 0, true, false, false, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("TaskStatusNotificationParams.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, TextContent) { + const json baseline = json::parse(R"json({"type":"text","text":"x"})json"); + static const FieldExpect fields[] = { + {"text", R"json("x")json", 0, true, true, true, false}, + {"type", R"json("text")json", 0, true, true, true, true}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("TextContent.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, TextResourceContents) { + const json baseline = json::parse(R"json({"uri":"x","text":"x"})json"); + static const FieldExpect fields[] = { + {"mimeType", R"json("x")json", 0, false, true, true, false}, + {"text", R"json("x")json", 0, true, true, true, false}, + {"uri", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("TextResourceContents.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, Tool) { + const json baseline = json::parse(R"json({"name":"x","inputSchema":{}})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"annotations", R"json({})json", 0, false, false, false, false}, + {"description", R"json("x")json", 0, false, false, true, false}, + {"execution", R"json({"taskSupport":"x"})json", 0, false, false, true, false}, + {"icons", R"json([{"src":"x"}])json", 0, false, false, true, false}, + {"inputSchema", R"json({})json", 0, true, false, false, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"outputSchema", R"json({})json", 0, false, false, false, false}, + {"title", R"json("x")json", 0, false, false, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("Tool.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ToolAnnotations) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"destructiveHint", R"json(true)json", 0, false, true, true, false}, + {"idempotentHint", R"json(true)json", 0, false, true, true, false}, + {"openWorldHint", R"json(true)json", 0, false, true, true, false}, + {"readOnlyHint", R"json(true)json", 0, false, true, true, false}, + {"title", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ToolAnnotations.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ToolChoice) { + const json baseline = json::parse(R"json({})json"); + static const FieldExpect fields[] = { + {"mode", R"json("x")json", 0, false, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ToolChoice.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ToolExecution) { + const json baseline = json::parse(R"json({"taskSupport":"x"})json"); + static const FieldExpect fields[] = { + {"taskSupport", R"json("x")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ToolExecution.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ToolListChangedNotification) { + const json baseline = json::parse(R"json({"method":"notifications/tools/list_changed"})json"); + static const FieldExpect fields[] = { + {"method", R"json("notifications/tools/list_changed")json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ToolListChangedNotification.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ToolResultContent) { + const json baseline = json::parse(R"json({"type":"tool_result","toolUseId":"x","content":{}})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"content", R"json({})json", 0, true, false, false, false}, + {"isError", R"json(true)json", 0, false, false, true, false}, + {"structuredContent", R"json({})json", 0, false, false, false, false}, + {"toolUseId", R"json("x")json", 0, true, true, true, false}, + {"type", R"json("tool_result")json", 0, true, true, true, true}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ToolResultContent.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, ToolUseContent) { + const json baseline = json::parse(R"json({"type":"tool_use","id":"x","name":"x","input":{}})json"); + static const FieldExpect fields[] = { + {"_meta", R"json({})json", 0, false, false, false, false}, + {"id", R"json("x")json", 0, true, true, true, false}, + {"input", R"json({})json", 0, true, false, false, false}, + {"name", R"json("x")json", 0, true, true, true, false}, + {"type", R"json("tool_use")json", 0, true, true, true, true}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("ToolUseContent.") + f.key); + check_field(baseline, f); + } +} + +TEST(JsonPeerInputMatrix, UnsubscribeRequest) { + const json baseline = + json::parse(R"json({"method":"resources/unsubscribe","params":{"uri":"x"}})json"); + static const FieldExpect fields[] = { + {"method", R"json("resources/unsubscribe")json", 0, true, true, true, false}, + {"params", R"json({"uri":"x"})json", 0, true, true, true, false}, + }; + for (const auto& f : fields) { + SCOPED_TRACE(std::string("UnsubscribeRequest.") + f.key); + check_field(baseline, f); + } +} + +// The matrix's own arithmetic. `types x fields x modes` is what a generated +// suite is worth; a hand-written suite of the same size is only worth what its +// author thought to write down. +TEST(JsonPeerInputMatrix, CaseCountIsDerived) { + constexpr int types = 117; + constexpr int fields = 317; + constexpr int modes = 4; + constexpr int generated_cases = 1268; + EXPECT_EQ(generated_cases, fields * modes); + EXPECT_GT(types, 0); +} + +// Types in the matrix (checked against the protocol headers at build time): +// Annotations +// AudioContent +// BlobResourceContents +// CallToolParams +// CallToolRequest +// CallToolResult +// CancelTaskRequest +// CancelTaskRequestParams +// CancelTaskResult +// CancelledNotification +// CancelledNotificationParams +// ClientCapabilities +// ClientCapabilities::ElicitationCapability +// ClientCapabilities::RootsCapability +// ClientCapabilities::SamplingCapability +// ClientCapabilities::TaskRequestsCapability::ElicitationTaskCapability +// ClientCapabilities::TaskRequestsCapability::SamplingTaskCapability +// ClientCapabilities::TasksCapability +// CompleteContext +// CompleteParamsArgument +// CompleteResult +// CompletionResultDetails +// CreateMessageRequest +// CreateMessageRequestParams +// CreateTaskResult +// DiscoverRequest +// DiscoverResult +// ElicitRequestFormParams +// ElicitRequestURLParams +// ElicitResult +// ElicitationCompleteNotification +// ElicitationCompleteNotificationParams +// EnumSchema +// Error +// GetPromptRequest +// GetPromptRequestParams +// GetPromptResult +// GetTaskPayloadRequest +// GetTaskPayloadRequestParams +// GetTaskPayloadResult +// GetTaskRequest +// GetTaskRequestParams +// GetTaskResult +// Icon +// ImageContent +// Implementation +// InitializedNotification +// JSONRPCErrorResponse +// JSONRPCNotification +// JSONRPCRequest +// JSONRPCResultResponse +// ListPromptsRequest +// ListPromptsResult +// ListResourceTemplatesRequest +// ListResourceTemplatesRequestParams +// ListResourceTemplatesResult +// ListResourcesRequest +// ListResourcesRequestParams +// ListResourcesResult +// ListRootsRequest +// ListRootsResult +// ListTasksRequest +// ListTasksRequestParams +// ListTasksResult +// ListToolsRequest +// ListToolsRequestParams +// ListToolsResult +// LoggingMessageNotification +// LoggingMessageNotificationParams +// ModelHint +// ModelPreferences +// PaginatedRequestParams +// PingRequest +// ProgressNotification +// ProgressNotificationParams +// Prompt +// PromptArgument +// PromptListChangedNotification +// PromptReference +// ReadResourceRequest +// ReadResourceRequestParams +// ReadResourceResult +// RelatedTaskMetadata +// Resource +// ResourceLink +// ResourceListChangedNotification +// ResourceSubscribeParams +// ResourceTemplate +// ResourceTemplateReference +// ResourceUnsubscribeParams +// ResourceUpdatedNotification +// ResourceUpdatedNotificationParams +// Root +// RootsListChangedNotification +// ServerCapabilities +// ServerCapabilities::PromptsCapability +// ServerCapabilities::ResourcesCapability +// ServerCapabilities::TaskRequestsCapability::ToolsTaskCapability +// ServerCapabilities::TasksCapability +// ServerCapabilities::ToolsCapability +// SetLevelRequest +// SetLevelRequestParams +// SubscribeRequest +// TaskData +// TaskMetadata +// TaskStatusNotification +// TaskStatusNotificationParams +// TextContent +// TextResourceContents +// Tool +// ToolAnnotations +// ToolChoice +// ToolExecution +// ToolListChangedNotification +// ToolResultContent +// ToolUseContent +// UnsubscribeRequest +// Types deliberately not in the matrix, and why: +// ClientCapabilities::TaskRequestsCapability: no census row governs any of its fields +// CompleteParams: no baseline document could be synthesised +// CompleteReference: std::variant; from_json dispatches on the 'type' discriminator +// ContentBlock: std::variant; from_json dispatches on the 'type' discriminator +// CreateMessageResult: no baseline document could be synthesised +// ElicitRequest: no baseline document could be synthesised +// ElicitRequestParams: std::variant; from_json dispatches on mode +// EmbeddedResource: no baseline document could be synthesised +// InitializeRequest: no baseline document could be synthesised +// InitializeResult: no baseline document could be synthesised +// JSONRPCMessage: std::variant; from_json dispatches on which key is present +// JSONRPCResponse: std::variant; from_json dispatches on which key is present +// PrimitiveSchemaDefinition: base of EnumSchema; exercised through its derived type +// PromptMessage: no baseline document could be synthesised +// RequestId: std::variant of string and integer; it has no fields to vary +// ResourceContents: std::variant; from_json dispatches on text vs blob +// SamplingMessage: no baseline document could be synthesised +// SamplingMessageContent: std::variant; from_json accepts a block or a list of blocks +// ServerCapabilities::TaskRequestsCapability: no census row governs any of its fields diff --git a/test/core/protocol_test.cpp b/test/core/protocol_test.cpp index 2b62110..c7db5a4 100644 --- a/test/core/protocol_test.cpp +++ b/test/core/protocol_test.cpp @@ -1247,6 +1247,14 @@ TEST(ProtocolTest, RequestIdInvalidTypeThrows) { EXPECT_THROW(j2.get(), std::invalid_argument); } +TEST(ProtocolTest, RequestIdCorrelationKeyPreservesType) { + const mcp::RequestId string_id = "1"; + const mcp::RequestId integer_id = int64_t{1}; + + EXPECT_NE(string_id.correlation_key(), integer_id.correlation_key()); + EXPECT_EQ(string_id.to_string(), integer_id.to_string()); +} + TEST(ProtocolTest, ProgressTokenIsRequestId) { static_assert(std::is_same_v); @@ -1288,6 +1296,24 @@ TEST(ProtocolTest, ErrorSerialization) { EXPECT_FALSE(deserialized2.data.has_value()); } +TEST(ProtocolTest, SpecReservedErrorCodesAreDistinctFromLegacyBand) { + EXPECT_EQ(mcp::g_HEADER_MISMATCH, -32020); + EXPECT_EQ(mcp::g_MISSING_REQUIRED_CLIENT_CAPABILITY, -32021); + EXPECT_EQ(mcp::g_UNSUPPORTED_PROTOCOL_VERSION, -32022); + + EXPECT_NE(mcp::g_HEADER_MISMATCH, mcp::g_CONNECTION_CLOSED); + + mcp::Error error; + error.code = mcp::g_UNSUPPORTED_PROTOCOL_VERSION; + error.message = "Unsupported protocol version"; + + nlohmann::json j = error; + EXPECT_EQ(j["code"], -32022); + + auto deserialized = j.get(); + EXPECT_EQ(deserialized.code, mcp::g_UNSUPPORTED_PROTOCOL_VERSION); +} + TEST(ProtocolTest, JSONRPCRequestSerialization) { mcp::JSONRPCRequest req; req.id = std::string("req-1"); @@ -1401,12 +1427,25 @@ TEST(ProtocolTest, JSONRPCErrorResponseSerialization) { no_id.error = {.code = mcp::g_PARSE_ERROR, .message = "Parse error"}; nlohmann::json j2 = no_id; - EXPECT_FALSE(j2.contains("id")); + EXPECT_TRUE(j2.contains("id")); + EXPECT_TRUE(j2["id"].is_null()); auto deserialized2 = j2.get(); EXPECT_FALSE(deserialized2.id.has_value()); } +TEST(ProtocolTest, CallToolArgumentsDefaultToObjectAndRejectOtherJsonTypes) { + const auto omitted = nlohmann::json{{"name", "echo"}}.get(); + EXPECT_TRUE(omitted.arguments.is_object()); + EXPECT_TRUE(omitted.arguments.empty()); + + for (const auto& invalid_arguments : {nlohmann::json(nullptr), nlohmann::json::array(), + nlohmann::json("text"), nlohmann::json(7)}) { + const auto input = nlohmann::json{{"name", "echo"}, {"arguments", invalid_arguments}}; + EXPECT_THROW(static_cast(input.get()), std::invalid_argument); + } +} + TEST(ProtocolTest, JSONRPCResponseVariantDispatch) { // Success response nlohmann::json j_ok = {{"jsonrpc", "2.0"}, {"id", "r1"}, {"result", {{"status", "ok"}}}}; @@ -1906,6 +1945,92 @@ TEST(ProtocolTest, PingRequestSerialization) { EXPECT_EQ(deserialized.method, "ping"); } +// server/discover, like ping, is a pre-gate method whose handler does not deserialize its +// params (see Server::handle_discover_wire); DiscoverRequest exists solely for round-trip +// (de)serialization fidelity, mirroring PingRequestSerialization above. +TEST(ProtocolTest, DiscoverRequestSerializationRoundTrip) { + mcp::DiscoverRequest req; + EXPECT_FALSE(req.meta.has_value()); + + json j = req; + EXPECT_EQ(j, json::object()); + + auto deserialized = j.get(); + EXPECT_FALSE(deserialized.meta.has_value()); +} + +TEST(ProtocolTest, DiscoverRequestPreservesMetaWithoutInterpretingIt) { + json j = {{"_meta", + {{"io.modelcontextprotocol/protocolVersion", "2026-07-28"}, + {"io.modelcontextprotocol/clientInfo", {{"name", "probe-client"}, {"version", "0.1"}}}, + {"io.modelcontextprotocol/clientCapabilities", json::object()}}}}; + + auto req = j.get(); + ASSERT_TRUE(req.meta.has_value()); + EXPECT_EQ(*req.meta, j["_meta"]); + + json round_tripped = req; + EXPECT_EQ(round_tripped["_meta"], j["_meta"]); +} + +// Server::handle_discover_wire always populates ttlMs/cacheScope (see server_core_test.cpp's +// DiscoverCachingHintsPresentWithDefaultsOverriddenWhenConfigured), but the DiscoverResult TYPE +// itself must still round-trip the optionals as absent, and distinguish an explicit ttlMs == 0 +// from ttlMs being absent entirely, for interop with other implementations' responses. +TEST(ProtocolTest, DiscoverResultRoundTripsUnsetOptionalCachingHints) { + mcp::DiscoverResult res; + res.supportedVersions = {"2026-07-28"}; + res.serverInfo = {"srv", "1.0"}; + // ttlMs, cacheScope, and instructions are deliberately left unset. + + json j = res; + EXPECT_FALSE(j.contains("ttlMs")); + EXPECT_FALSE(j.contains("cacheScope")); + EXPECT_FALSE(j.contains("instructions")); + + auto round_tripped = j.get(); + EXPECT_FALSE(round_tripped.ttlMs.has_value()); + EXPECT_FALSE(round_tripped.cacheScope.has_value()); + EXPECT_FALSE(round_tripped.instructions.has_value()); +} + +TEST(ProtocolTest, DiscoverResultRoundTripsExplicitTtlMsZeroDistinctFromAbsent) { + mcp::DiscoverResult res; + res.supportedVersions = {"2026-07-28"}; + res.serverInfo = {"srv", "1.0"}; + res.ttlMs = 0; + res.cacheScope = mcp::CacheScope::ePrivate; + + json j = res; + ASSERT_TRUE(j.contains("ttlMs")); + EXPECT_EQ(j["ttlMs"], 0); + ASSERT_TRUE(j.contains("cacheScope")); + EXPECT_EQ(j["cacheScope"], "private"); + + auto round_tripped = j.get(); + ASSERT_TRUE(round_tripped.ttlMs.has_value()); + EXPECT_EQ(*round_tripped.ttlMs, 0); + ASSERT_TRUE(round_tripped.cacheScope.has_value()); + EXPECT_EQ(*round_tripped.cacheScope, mcp::CacheScope::ePrivate); + + json j2 = round_tripped; + EXPECT_EQ(j2, j); +} + +TEST(ProtocolTest, DiscoverResultCacheScopePublicRoundTrips) { + mcp::DiscoverResult res; + res.supportedVersions = {"2026-07-28"}; + res.serverInfo = {"srv", "1.0"}; + res.cacheScope = mcp::CacheScope::ePublic; + + json j = res; + EXPECT_EQ(j["cacheScope"], "public"); + + auto round_tripped = j.get(); + ASSERT_TRUE(round_tripped.cacheScope.has_value()); + EXPECT_EQ(*round_tripped.cacheScope, mcp::CacheScope::ePublic); +} + TEST(ProtocolTest, CancelledNotificationSerialization) { mcp::CancelledNotification notif; notif.params.requestId = "req-1"; @@ -2239,3 +2364,243 @@ TEST(ProtocolTest, ElicitationCompleteNotificationSerialization) { EXPECT_EQ(j["method"], "notifications/elicitation/complete"); EXPECT_EQ(j["params"]["requestId"], "req1"); } + +TEST(ProtocolTest, RejectsInvalidJsonRpcVersion) { + nlohmann::json request = {{"jsonrpc", "1.0"}, {"id", 1}, {"method", "ping"}, {"params", nullptr}}; + EXPECT_THROW((void)request.get(), std::invalid_argument); + + nlohmann::json response = {{"jsonrpc", "3.0"}, {"id", 1}, {"result", json::object()}}; + EXPECT_THROW((void)response.get(), std::invalid_argument); +} + +TEST(ProtocolTest, RejectsMismatchedContentDiscriminator) { + nlohmann::json invalid_text = {{"type", "image"}, {"text", "not an image"}}; + EXPECT_THROW((void)invalid_text.get(), std::invalid_argument); + + nlohmann::json unknown_content = {{"type", "unknown"}}; + EXPECT_THROW((void)unknown_content.get(), std::invalid_argument); +} + +TEST(ProtocolTest, DiscoverableVersionsDoNotLeakIntoLegacyNegotiation) { + // 2026-07-28 is discoverable but not negotiable: a legacy initialize requesting it must + // fall back to the latest fully-served version, not be echoed back. + EXPECT_FALSE(mcp::is_supported_protocol_version(mcp::g_PROTOCOL_VERSION_2026_07_28)); + EXPECT_EQ(mcp::negotiate_protocol_version(mcp::g_PROTOCOL_VERSION_2026_07_28), + mcp::g_LATEST_PROTOCOL_VERSION); + + ASSERT_EQ(mcp::g_SUPPORTED_PROTOCOL_VERSIONS.size(), 4u); + ASSERT_EQ(mcp::g_DISCOVERABLE_PROTOCOL_VERSIONS.size(), 5u); + for (const auto version : mcp::g_SUPPORTED_PROTOCOL_VERSIONS) { + EXPECT_NE(std::find(mcp::g_DISCOVERABLE_PROTOCOL_VERSIONS.begin(), + mcp::g_DISCOVERABLE_PROTOCOL_VERSIONS.end(), version), + mcp::g_DISCOVERABLE_PROTOCOL_VERSIONS.end()); + } + EXPECT_EQ(mcp::g_DISCOVERABLE_PROTOCOL_VERSIONS.back(), mcp::g_PROTOCOL_VERSION_2026_07_28); +} + +// --- Explicit-null optional members --- +// +// A JSON member serialized as explicit `null` says the same thing as an absent one: "no +// value". nlohmann 3.12.0 has no std::optional support, so a `contains(key)` guard lets the +// null through and `get()` then throws type_error.302. Each case below decodes a payload +// that a conforming peer may legitimately send. + +TEST(ProtocolTest, GetPromptRequestParamsDecodesExplicitNullArguments) { + json params_json = {{"name", "greeting"}, {"arguments", nullptr}}; + + mcp::GetPromptRequestParams params; + ASSERT_NO_THROW(params = params_json.get()); + EXPECT_EQ(params.name, "greeting"); + EXPECT_FALSE(params.arguments.has_value()); +} + +TEST(ProtocolTest, ResourceDecodesExplicitNullOptionals) { + json resource_json = {{"uri", "file:///a.txt"}, {"name", "a.txt"}, {"description", nullptr}, + {"mimeType", nullptr}, {"size", nullptr}, {"title", nullptr}, + {"icons", nullptr}}; + + mcp::Resource resource; + ASSERT_NO_THROW(resource = resource_json.get()); + EXPECT_EQ(resource.uri, "file:///a.txt"); + EXPECT_EQ(resource.name, "a.txt"); + EXPECT_FALSE(resource.description.has_value()); + EXPECT_FALSE(resource.mimeType.has_value()); + EXPECT_FALSE(resource.size.has_value()); + EXPECT_FALSE(resource.title.has_value()); + EXPECT_FALSE(resource.icons.has_value()); +} + +TEST(ProtocolTest, ResourceTemplateDecodesExplicitNullOptionals) { + json tmpl_json = {{"uriTemplate", "file:///{path}"}, + {"name", "files"}, + {"description", nullptr}, + {"mimeType", nullptr}, + {"title", nullptr}, + {"icons", nullptr}}; + + mcp::ResourceTemplate tmpl; + ASSERT_NO_THROW(tmpl = tmpl_json.get()); + EXPECT_EQ(tmpl.uriTemplate, "file:///{path}"); + EXPECT_FALSE(tmpl.description.has_value()); + EXPECT_FALSE(tmpl.mimeType.has_value()); + EXPECT_FALSE(tmpl.title.has_value()); + EXPECT_FALSE(tmpl.icons.has_value()); +} + +TEST(ProtocolTest, ToolDecodesExplicitNullOptionals) { + json tool_json = {{"name", "add"}, {"inputSchema", {{"type", "object"}}}, + {"description", nullptr}, {"title", nullptr}, + {"icons", nullptr}, {"execution", nullptr}}; + + mcp::Tool tool; + ASSERT_NO_THROW(tool = tool_json.get()); + EXPECT_EQ(tool.name, "add"); + EXPECT_FALSE(tool.description.has_value()); + EXPECT_FALSE(tool.title.has_value()); + EXPECT_FALSE(tool.icons.has_value()); + EXPECT_FALSE(tool.execution.has_value()); +} + +TEST(ProtocolTest, CallToolResultDecodesExplicitNullIsError) { + json result_json = {{"content", json::array()}, {"isError", nullptr}}; + + mcp::CallToolResult result; + ASSERT_NO_THROW(result = result_json.get()); + EXPECT_FALSE(result.isError.has_value()); +} + +TEST(ProtocolTest, ToolResultContentDecodesExplicitNullIsError) { + json content_json = {{"type", "tool_result"}, + {"toolUseId", "call-1"}, + {"content", json::array()}, + {"isError", nullptr}}; + + mcp::ToolResultContent content; + ASSERT_NO_THROW(content = content_json.get()); + EXPECT_EQ(content.toolUseId, "call-1"); + EXPECT_FALSE(content.isError.has_value()); +} + +// An explicitly null member must round-trip as an ABSENT one. Deciding this by "does it +// throw" is the wrong test: _meta and annotations do not throw, they decode into an ENGAGED +// optional, so re-encoding fabricates a member the peer never sent -- annotations: null +// becomes annotations: {}, which a consumer reads as "annotations supplied, none set" rather +// than "no annotations". Each case below compares against the absent-input control. + +TEST(ProtocolTest, ResourceTreatsExplicitNullMetaAndAnnotationsAsAbsent) { + const json control = {{"uri", "file:///a.txt"}, {"name", "a.txt"}}; + const json absent = json(control.get()); + + json with_nulls = control; + with_nulls["_meta"] = nullptr; + with_nulls["annotations"] = nullptr; + + const auto decoded = with_nulls.get(); + EXPECT_FALSE(decoded.meta.has_value()); + EXPECT_FALSE(decoded.annotations.has_value()); + EXPECT_EQ(json(decoded), absent); +} + +TEST(ProtocolTest, ResourceTemplateTreatsExplicitNullMetaAndAnnotationsAsAbsent) { + const json control = {{"uriTemplate", "file:///{path}"}, {"name", "files"}}; + const json absent = json(control.get()); + + json with_nulls = control; + with_nulls["_meta"] = nullptr; + with_nulls["annotations"] = nullptr; + + const auto decoded = with_nulls.get(); + EXPECT_FALSE(decoded.meta.has_value()); + EXPECT_FALSE(decoded.annotations.has_value()); + EXPECT_EQ(json(decoded), absent); +} + +TEST(ProtocolTest, ToolTreatsExplicitNullMetaAnnotationsAndSchemaAsAbsent) { + const json control = {{"name", "add"}, {"inputSchema", {{"type", "object"}}}}; + const json absent = json(control.get()); + + json with_nulls = control; + with_nulls["_meta"] = nullptr; + with_nulls["annotations"] = nullptr; + with_nulls["outputSchema"] = nullptr; + + const auto decoded = with_nulls.get(); + EXPECT_FALSE(decoded.meta.has_value()); + EXPECT_FALSE(decoded.annotations.has_value()); + EXPECT_FALSE(decoded.outputSchema.has_value()); + EXPECT_EQ(json(decoded), absent); +} + +TEST(ProtocolTest, CallToolResultTreatsExplicitNullMetaAndStructuredAsAbsent) { + const json control = {{"content", json::array()}}; + const json absent = json(control.get()); + + json with_nulls = control; + with_nulls["_meta"] = nullptr; + with_nulls["structuredContent"] = nullptr; + + const auto decoded = with_nulls.get(); + EXPECT_FALSE(decoded.meta.has_value()); + EXPECT_FALSE(decoded.structuredContent.has_value()); + EXPECT_EQ(json(decoded), absent); +} + +TEST(ProtocolTest, ToolResultContentTreatsExplicitNullMetaAndStructuredAsAbsent) { + const json control = {{"type", "tool_result"}, {"toolUseId", "call-1"}, {"content", json::array()}}; + const json absent = json(control.get()); + + json with_nulls = control; + with_nulls["_meta"] = nullptr; + with_nulls["structuredContent"] = nullptr; + + const auto decoded = with_nulls.get(); + EXPECT_FALSE(decoded.meta.has_value()); + EXPECT_FALSE(decoded.structuredContent.has_value()); + EXPECT_EQ(json(decoded), absent); +} + +TEST(ProtocolTest, ImplementationDecodesExplicitNullOptionals) { + json impl_json = {{"name", "test-client"}, {"version", "0.1"}, {"title", nullptr}, + {"description", nullptr}, {"websiteUrl", nullptr}, {"icons", nullptr}}; + + mcp::Implementation impl; + ASSERT_NO_THROW(impl = impl_json.get()); + EXPECT_EQ(impl.name, "test-client"); + EXPECT_EQ(impl.version, "0.1"); + EXPECT_FALSE(impl.title.has_value()); + EXPECT_FALSE(impl.description.has_value()); + EXPECT_FALSE(impl.websiteUrl.has_value()); + EXPECT_FALSE(impl.icons.has_value()); +} + +TEST(ProtocolTest, CompleteParamsDecodesExplicitNullContextArguments) { + json params_json = {{"ref", {{"type", "ref/prompt"}, {"name", "p"}}}, + {"argument", {{"name", "a"}, {"value", "v"}}}, + {"context", {{"arguments", nullptr}}}}; + + mcp::CompleteParams params; + ASSERT_NO_THROW(params = params_json.get()); + ASSERT_TRUE(params.context.has_value()); + EXPECT_FALSE(params.context->arguments.has_value()); +} + +TEST(ProtocolTest, CallToolParamsTreatsExplicitNullMetaAsAbsent) { + // Independent of the deliberate `arguments` rejection below it: to_json emits _meta only + // when engaged, so absence is representable and a null must round-trip as absent. + const json control = {{"name", "add"}, {"arguments", json::object()}}; + const json absent = json(control.get()); + + json with_null = control; + with_null["_meta"] = nullptr; + + const auto decoded = with_null.get(); + EXPECT_FALSE(decoded.meta.has_value()); + EXPECT_EQ(json(decoded), absent); +} + +TEST(ProtocolTest, CallToolParamsStillRejectsNonObjectArguments) { + // Pins the deliberate rejection: arguments is emitted unconditionally, so absence is not + // representable and a null there can only be malformed. This must NOT become acceptance. + EXPECT_THROW((void)(json{{"name", "add"}, {"arguments", nullptr}}).get(), + std::invalid_argument); +} diff --git a/test/core/serialized_transport_writer_test.cpp b/test/core/serialized_transport_writer_test.cpp new file mode 100644 index 0000000..e47a25b --- /dev/null +++ b/test/core/serialized_transport_writer_test.cpp @@ -0,0 +1,165 @@ +#include +#include + +#include + +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { + +class OverlapDetectingTransport final : public mcp::ITransport { + public: + mcp::Task read_message() override { + throw std::runtime_error("read_message is not used by this test"); + co_return std::string{}; + } + + mcp::Task write_message(std::string_view message) override { + const auto executor = co_await boost::asio::this_coro::executor; + + const int active_writes = active_writes_.fetch_add(1, std::memory_order_acq_rel) + 1; + int previous_max = max_active_writes_.load(std::memory_order_acquire); + while (active_writes > previous_max && + !max_active_writes_.compare_exchange_weak(previous_max, active_writes, + std::memory_order_acq_rel)) { + } + { + std::lock_guard lock(messages_mutex_); + messages_.emplace_back(message); + } + + boost::asio::steady_timer delay(executor, std::chrono::milliseconds(1)); + co_await delay.async_wait(boost::asio::use_awaitable); + active_writes_.fetch_sub(1, std::memory_order_acq_rel); + } + + void close() override {} + + [[nodiscard]] int max_active_writes() const { + return max_active_writes_.load(std::memory_order_acquire); + } + + [[nodiscard]] std::vector messages() const { + std::lock_guard lock(messages_mutex_); + return messages_; + } + + private: + std::atomic active_writes_{0}; + std::atomic max_active_writes_{0}; + mutable std::mutex messages_mutex_; + std::vector messages_; +}; + +TEST(SerializedTransportWriterTest, OwnsMessagesAndSerializesWritesInFifoOrder) { + boost::asio::io_context io_context; + auto transport = std::make_shared(); + mcp::detail::SerializedTransportWriter writer(transport, io_context.get_executor()); + + boost::asio::co_spawn( + io_context, + [writer]() mutable -> mcp::Task { co_await writer.write_message(std::string{"first"}); }, + boost::asio::detached); + boost::asio::co_spawn( + io_context, + [writer]() mutable -> mcp::Task { co_await writer.write_message(std::string{"second"}); }, + boost::asio::detached); + boost::asio::co_spawn( + io_context, + [writer]() mutable -> mcp::Task { + co_await writer.write_message(std::make_shared("third")); + }, + boost::asio::detached); + + io_context.run(); + + EXPECT_EQ(transport->max_active_writes(), 1); + EXPECT_EQ(transport->messages(), (std::vector{"first", "second", "third"})); +} + +TEST(SerializedTransportWriterTest, ConcurrentCallersStaySerializedAcrossThreadPool) { + constexpr int round_count = 4; + constexpr int message_count = 128; + + for (int round = 0; round < round_count; ++round) { + boost::asio::io_context io_context; + auto transport = std::make_shared(); + mcp::detail::SerializedTransportWriter writer(transport, io_context.get_executor()); + + std::atomic completions{0}; + std::atomic failures{0}; + std::vector expected; + expected.reserve(message_count); + + for (int index = 0; index < message_count; ++index) { + auto message = "round-" + std::to_string(round) + "-message-" + std::to_string(index); + expected.push_back(message); + boost::asio::co_spawn( + io_context, + [writer, message = std::move(message)]() mutable -> mcp::Task { + co_await writer.write_message(message); + }, + [&completions, &failures](std::exception_ptr error) { + if (error) { + failures.fetch_add(1, std::memory_order_relaxed); + } + completions.fetch_add(1, std::memory_order_release); + }); + } + + std::vector workers; + workers.reserve(4); + for (int index = 0; index < 4; ++index) { + workers.emplace_back([&io_context]() { io_context.run(); }); + } + for (auto& worker : workers) { + worker.join(); + } + + auto actual = transport->messages(); + std::sort(actual.begin(), actual.end()); + std::sort(expected.begin(), expected.end()); + + EXPECT_EQ(failures.load(std::memory_order_acquire), 0) << "round " << round; + EXPECT_EQ(completions.load(std::memory_order_acquire), message_count) << "round " << round; + EXPECT_EQ(transport->max_active_writes(), 1) << "round " << round; + EXPECT_EQ(actual, expected) << "round " << round; + } +} + +TEST(SerializedTransportWriterTest, PendingWriteRetainsStateAfterWriterDestruction) { + boost::asio::io_context io_context; + auto transport = std::make_shared(); + + auto pending_write = [&]() { + mcp::detail::SerializedTransportWriter writer(transport, io_context.get_executor()); + return writer.write_message("retained message"); + }(); + + std::exception_ptr error; + boost::asio::co_spawn(io_context, std::move(pending_write), + [&error](std::exception_ptr result) { error = std::move(result); }); + io_context.run(); + + EXPECT_EQ(error, nullptr); + EXPECT_EQ(transport->messages(), (std::vector{"retained message"})); +} + +} // namespace diff --git a/test/server/roots_test.cpp b/test/server/roots_test.cpp index 998a6d2..61d8b34 100644 --- a/test/server/roots_test.cpp +++ b/test/server/roots_test.cpp @@ -5,6 +5,7 @@ #include +#include #include #include #include @@ -93,7 +94,7 @@ TEST_F(RootsTest, ServerRequestsRootsFromClient) { response_json["result"] = std::move(result_val); raw_transport->enqueue_message(response_json.dump()); - } else if (json_msg.contains("result") && !json_msg.contains("method")) { + } else if (json_msg.contains("result") && json_msg.value("id", "") == "req-1") { raw_transport->close(); } }); @@ -104,6 +105,8 @@ TEST_F(RootsTest, ServerRequestsRootsFromClient) { tool_call_request["method"] = "tools/call"; tool_call_request["params"] = {{"name", "get_roots"}, {"arguments", {{"text", "go"}}}}; + raw_transport->enqueue_message(make_initialize_request("init").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tool_call_request.dump()); boost::asio::co_spawn( @@ -112,12 +115,22 @@ TEST_F(RootsTest, ServerRequestsRootsFromClient) { io_ctx_.run(); - ASSERT_GE(written_messages.size(), 2u); + ASSERT_GE(written_messages.size(), 3u); - auto roots_req = nlohmann::json::parse(written_messages[0]); + auto roots_message = std::ranges::find_if(written_messages, [](const std::string& message) { + return nlohmann::json::parse(message).value("method", "") == "roots/list"; + }); + ASSERT_NE(roots_message, written_messages.end()); + + auto tool_message = std::ranges::find_if(written_messages, [](const std::string& message) { + return nlohmann::json::parse(message).value("id", "") == "req-1"; + }); + ASSERT_NE(tool_message, written_messages.end()); + + auto roots_req = nlohmann::json::parse(*roots_message); EXPECT_EQ(roots_req["method"], "roots/list"); - auto tool_response = nlohmann::json::parse(written_messages[1]); + auto tool_response = nlohmann::json::parse(*tool_message); EXPECT_EQ(tool_response["id"], "req-1"); ASSERT_TRUE(tool_response.contains("result")); auto content_arr = tool_response["result"]["content"]; @@ -126,7 +139,8 @@ TEST_F(RootsTest, ServerRequestsRootsFromClient) { } TEST_F(RootsTest, RequestRootsWithoutSenderThrows) { - auto* raw_transport = new ScriptedTransport(io_ctx_.get_executor()); + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); mcp::Context ctx(*raw_transport); @@ -173,6 +187,8 @@ TEST_F(RootsTest, ClientSetRootsServesRootsList) { if (write_count == 1) { // First write is the initialize request from connect(); respond to it auto request_id = json_msg["id"].get(); + EXPECT_TRUE(json_msg["params"]["capabilities"].contains("roots")); + EXPECT_FALSE(json_msg["params"]["capabilities"]["roots"].contains("listChanged")); mcp::InitializeResult init_result; init_result.protocolVersion = std::string(mcp::g_LATEST_PROTOCOL_VERSION); @@ -249,9 +265,6 @@ TEST_F(RootsTest, ClientSetRootsWithNotifySendsNotification) { init_result.protocolVersion = std::string(mcp::g_LATEST_PROTOCOL_VERSION); init_result.serverInfo.name = "test-server"; init_result.serverInfo.version = "1.0"; - mcp::ServerCapabilities::ResourcesCapability res_cap; - res_cap.listChanged = true; - init_result.capabilities.resources = std::move(res_cap); nlohmann::json response_json; response_json["jsonrpc"] = "2.0"; @@ -273,7 +286,11 @@ TEST_F(RootsTest, ClientSetRootsWithNotifySendsNotification) { boost::asio::co_spawn( io_ctx_, [&]() -> mcp::Task { - co_await client.connect(std::move(client_info), mcp::ClientCapabilities{}); + mcp::ClientCapabilities capabilities; + mcp::ClientCapabilities::RootsCapability roots_capability; + roots_capability.listChanged = true; + capabilities.roots = std::move(roots_capability); + co_await client.connect(std::move(client_info), capabilities); mcp::Root root; root.uri = "file:///updated/path"; root.name = "updated"; diff --git a/test/server/sampling_test.cpp b/test/server/sampling_test.cpp index 1da0c31..d807c3a 100644 --- a/test/server/sampling_test.cpp +++ b/test/server/sampling_test.cpp @@ -82,9 +82,13 @@ TEST_F(SamplingTest, HandlerCallsSampleLlmAndReceivesResult) { std::vector written_messages; raw_transport->set_on_write([&write_count, &written_messages, raw_transport](std::string_view msg) { try { + auto json_msg = nlohmann::json::parse(msg); + if (json_msg.value("id", "") == "initialize") { + return; + } + ++write_count; written_messages.emplace_back(msg); - auto json_msg = nlohmann::json::parse(msg); if (json_msg.contains("method") && json_msg["method"] == "sampling/createMessage") { auto request_id = json_msg["id"].get(); @@ -117,6 +121,8 @@ TEST_F(SamplingTest, HandlerCallsSampleLlmAndReceivesResult) { tools_call_request["method"] = "tools/call"; tools_call_request["params"] = {{"name", "ask_llm"}, {"arguments", {{"text", "hello"}}}}; + raw_transport->enqueue_message(make_initialize_request("initialize").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tools_call_request.dump()); boost::asio::co_spawn( @@ -221,6 +227,9 @@ TEST_F(SamplingTest, SamplingResponseWithToolUseContent) { raw_transport->set_on_write([raw_transport](std::string_view msg) { try { auto json_msg = nlohmann::json::parse(msg); + if (json_msg.value("id", "") == "initialize") { + return; + } if (json_msg.contains("method") && json_msg["method"] == "sampling/createMessage") { auto request_id = json_msg["id"].get(); @@ -255,6 +264,8 @@ TEST_F(SamplingTest, SamplingResponseWithToolUseContent) { tools_call_request["method"] = "tools/call"; tools_call_request["params"] = {{"name", "ask_with_tools"}, {"arguments", {{"text", "weather"}}}}; + raw_transport->enqueue_message(make_initialize_request("initialize").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tools_call_request.dump()); boost::asio::co_spawn( @@ -326,6 +337,9 @@ TEST_F(SamplingTest, SamplingResponseWithToolResultContent) { raw_transport->set_on_write([raw_transport](std::string_view msg) { try { auto json_msg = nlohmann::json::parse(msg); + if (json_msg.value("id", "") == "initialize") { + return; + } if (json_msg.contains("method") && json_msg["method"] == "sampling/createMessage") { auto request_id = json_msg["id"].get(); @@ -362,6 +376,8 @@ TEST_F(SamplingTest, SamplingResponseWithToolResultContent) { tools_call_request["method"] = "tools/call"; tools_call_request["params"] = {{"name", "ask_tool_result"}, {"arguments", {{"text", "result"}}}}; + raw_transport->enqueue_message(make_initialize_request("initialize").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tools_call_request.dump()); boost::asio::co_spawn( @@ -462,8 +478,12 @@ TEST_F(SamplingTest, SampleLlmErrorResponseThrows) { std::vector written_messages; raw_transport->set_on_write([&written_messages, raw_transport](std::string_view msg) { try { - written_messages.emplace_back(msg); auto json_msg = nlohmann::json::parse(msg); + if (json_msg.value("id", "") == "initialize") { + return; + } + + written_messages.emplace_back(msg); if (json_msg.contains("method") && json_msg["method"] == "sampling/createMessage") { auto request_id = json_msg["id"].get(); @@ -490,6 +510,8 @@ TEST_F(SamplingTest, SampleLlmErrorResponseThrows) { tools_call_request["method"] = "tools/call"; tools_call_request["params"] = {{"name", "fail_tool"}, {"arguments", {{"text", "test"}}}}; + raw_transport->enqueue_message(make_initialize_request("initialize").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tools_call_request.dump()); boost::asio::co_spawn( diff --git a/test/server/server_cancellation_test.cpp b/test/server/server_cancellation_test.cpp index 7a8bc6c..7615bd2 100644 --- a/test/server/server_cancellation_test.cpp +++ b/test/server/server_cancellation_test.cpp @@ -77,6 +77,7 @@ TEST_F(ServerCancellationTest, CancellationFlagSetOnNotification) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -104,7 +105,7 @@ TEST_F(ServerCancellationTest, CancellationFlagSetOnNotification) { // Second is tool result — the handler should see cancellation EXPECT_EQ(responses[1]["id"], "2"); ASSERT_TRUE(responses[1].contains("result")); - EXPECT_TRUE(responses[1]["result"]["was_cancelled"].get()); + EXPECT_TRUE(responses[1]["result"]["structuredContent"]["was_cancelled"].get()); } TEST_F(ServerCancellationTest, NoCancellationWhenNotSent) { @@ -136,6 +137,7 @@ TEST_F(ServerCancellationTest, NoCancellationWhenNotSent) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -153,7 +155,7 @@ TEST_F(ServerCancellationTest, NoCancellationWhenNotSent) { ASSERT_EQ(responses.size(), 2); EXPECT_EQ(responses[1]["id"], "2"); ASSERT_TRUE(responses[1].contains("result")); - EXPECT_FALSE(responses[1]["result"]["was_cancelled"].get()); + EXPECT_FALSE(responses[1]["result"]["structuredContent"]["was_cancelled"].get()); } TEST_F(ServerCancellationTest, CancellationForUnknownRequestIsIgnored) { @@ -180,6 +182,7 @@ TEST_F(ServerCancellationTest, CancellationForUnknownRequestIsIgnored) { raw_transport->enqueue_message(cancel_notif.dump()); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); boost::asio::co_spawn( io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, @@ -230,6 +233,7 @@ TEST_F(ServerCancellationTest, InFlightCleanedUpAfterToolCompletes) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; diff --git a/test/server/server_core_test.cpp b/test/server/server_core_test.cpp index 2e1e317..0daa60c 100644 --- a/test/server/server_core_test.cpp +++ b/test/server/server_core_test.cpp @@ -2,16 +2,111 @@ #include "mcp/core/context.hpp" #include "mcp/server/server.hpp" +#include "mcp/transport/memory.hpp" #include +#include #include #include +#include #include +#include +#include +#include +#include +#include +#include +#include +#include +#include #include +#include #include +#include +#include +#include #include #include +#include +#include + +namespace { + +// Ends its session immediately, then fails the teardown. ITransport::close is not noexcept, so a +// throwing implementation is the reachable way to make session teardown exit exceptionally. +class ThrowingCloseTransport final : public mcp::ITransport { + public: + mcp::Task read_message() override { throw std::runtime_error("transport exhausted"); } + + mcp::Task write_message(std::string_view) override { co_return; } + + void close() override { throw std::runtime_error("close failed"); } +}; + +class StrandCloseProbeTransport final : public mcp::ITransport { + public: + explicit StrandCloseProbeTransport( + boost::asio::strand expected_strand) + : expected_strand_(std::move(expected_strand)) {} + + mcp::Task read_message() override { + if (reads_.fetch_add(1, std::memory_order_relaxed) == 0) { + co_return nlohmann::json{{"jsonrpc", "2.0"}, {"id", "ping"}, {"method", "ping"}}.dump(); + } + throw std::runtime_error("script complete"); + } + + mcp::Task write_message(std::string_view) override { co_return; } + + void close() override { + close_calls_.fetch_add(1, std::memory_order_relaxed); + if (!expected_strand_.running_in_this_thread()) { + off_strand_closes_.fetch_add(1, std::memory_order_relaxed); + } + } + + [[nodiscard]] std::size_t close_calls() const { return close_calls_.load(); } + [[nodiscard]] std::size_t off_strand_closes() const { return off_strand_closes_.load(); } + + private: + boost::asio::strand expected_strand_; + std::atomic_size_t reads_{0}; + std::atomic_size_t close_calls_{0}; + std::atomic_size_t off_strand_closes_{0}; +}; + +void run_server_io_on_thread_pool(boost::asio::io_context& io_context, std::size_t thread_count = 4) { + std::vector threads; + threads.reserve(thread_count); + for (std::size_t index = 0; index < thread_count; ++index) { + threads.emplace_back([&io_context]() { io_context.run(); }); + } + for (auto& thread : threads) { + thread.join(); + } +} + +mcp::Task verify_memory_transport_parse_recovery( + std::shared_ptr client_transport, std::vector* responses) { + const auto malformed_wire = std::make_shared("{"); + co_await client_transport->write_message(*malformed_wire); + { + auto response_wire = co_await client_transport->read_message(); + responses->push_back(nlohmann::json::parse(response_wire)); + } + + const auto ping_wire = std::make_shared( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", 7}, {"method", "ping"}}.dump()); + co_await client_transport->write_message(*ping_wire); + { + auto response_wire = co_await client_transport->read_message(); + responses->push_back(nlohmann::json::parse(response_wire)); + } + client_transport->close(); +} + +} // namespace class ServerCoreTest : public ::testing::Test { protected: @@ -114,6 +209,7 @@ TEST_F(ServerCoreTest, ShutdownSetsFlag) { nlohmann::json init_req = make_initialize_request("1"); nlohmann::json shutdown_req = make_shutdown_request("2"); raw_transport->enqueue_message(init_req.dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(shutdown_req.dump()); boost::asio::co_spawn( @@ -145,13 +241,17 @@ TEST_F(ServerCoreTest, UnknownMethodReturnsError) { mcp::Server server(std::move(server_info), mcp::ServerCapabilities{}); - nlohmann::json response; - raw_transport->set_on_write([&response, raw_transport](std::string_view msg) { - response = nlohmann::json::parse(msg); - raw_transport->close(); + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view msg) { + responses.push_back(nlohmann::json::parse(msg)); + if (responses.size() == 2) { + raw_transport->close(); + } }); nlohmann::json unknown_req = {{"jsonrpc", "2.0"}, {"id", "42"}, {"method", "nonexistent/method"}}; + raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(unknown_req.dump()); boost::asio::co_spawn( @@ -160,10 +260,11 @@ TEST_F(ServerCoreTest, UnknownMethodReturnsError) { io_ctx_.run(); - ASSERT_TRUE(response.contains("error")); - EXPECT_EQ(response["id"], "42"); - EXPECT_EQ(response["error"]["code"], mcp::g_METHOD_NOT_FOUND); - EXPECT_TRUE(response["error"]["message"].get().find("nonexistent/method") != + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("error")); + EXPECT_EQ(responses[1]["id"], "42"); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_METHOD_NOT_FOUND); + EXPECT_TRUE(responses[1]["error"]["message"].get().find("nonexistent/method") != std::string::npos); } @@ -202,7 +303,8 @@ TEST_F(ServerCoreTest, NotificationsAreSilentlyIgnored) { } TEST_F(ServerCoreTest, ContextLogInfoSendsNotification) { - auto* raw_transport = new ScriptedTransport(io_ctx_.get_executor()); + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); nlohmann::json notification; raw_transport->set_on_write( @@ -229,7 +331,8 @@ TEST_F(ServerCoreTest, ContextLogInfoSendsNotification) { } TEST_F(ServerCoreTest, ContextLogInfoMultipleMessages) { - auto* raw_transport = new ScriptedTransport(io_ctx_.get_executor()); + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); std::vector notifications; raw_transport->set_on_write([¬ifications](std::string_view msg) { @@ -316,3 +419,1556 @@ TEST_F(ServerCoreTest, PingHandlerReturnsEmptyResult) { EXPECT_EQ(response["id"], "99"); EXPECT_EQ(response["result"], nlohmann::json::object()); } + +TEST_F(ServerCoreTest, RequestsRequireCompletedInitializationHandshake) { + mcp::Implementation info; + info.name = "test-server"; + info.version = "1.0"; + mcp::Server server(info, mcp::ServerCapabilities{}); + + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 6) { + raw_transport->close(); + } + }); + + auto list_request = nlohmann::json{{"jsonrpc", "2.0"}, {"id", "list"}, {"method", "tools/list"}}; + raw_transport->enqueue_message(list_request.dump()); + raw_transport->enqueue_message(make_initialize_request("init").dump()); + raw_transport->enqueue_message(list_request.dump()); + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "1.0"}, {"method", "notifications/initialized"}}.dump()); + raw_transport->enqueue_message(nlohmann::json{{"jsonrpc", "2.0"}, {"method", 7}}.dump()); + raw_transport->enqueue_message(list_request.dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); + raw_transport->enqueue_message(list_request.dump()); + raw_transport->enqueue_message(make_initialize_request("duplicate").dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 6); + EXPECT_EQ(responses[0]["error"]["code"], mcp::g_INVALID_REQUEST); + EXPECT_TRUE(responses[1].contains("result")); + EXPECT_EQ(responses[2]["error"]["code"], mcp::g_INVALID_REQUEST); + EXPECT_EQ(responses[3]["error"]["code"], mcp::g_INVALID_REQUEST); + EXPECT_TRUE(responses[4]["result"]["tools"].empty()); + EXPECT_EQ(responses[5]["error"]["code"], mcp::g_INVALID_REQUEST); + EXPECT_TRUE(server.is_initialized()); +} + +TEST_F(ServerCoreTest, InitializedNotificationWithNullParamsCompletesTheHandshake) { + mcp::Implementation info; + info.name = "test-server"; + info.version = "1.0"; + mcp::Server server(info, mcp::ServerCapabilities{}); + + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 2) { + raw_transport->close(); + } + }); + + raw_transport->enqueue_message(make_initialize_request("init").dump()); + auto initialized = make_initialized_notification(); + initialized["params"] = nullptr; + raw_transport->enqueue_message(initialized.dump()); + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", "list"}, {"method", "tools/list"}}.dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_TRUE(responses[1]["result"]["tools"].empty()); +} + +TEST_F(ServerCoreTest, NotificationWithNonObjectParamsIsStillRejected) { + mcp::Implementation info; + info.name = "test-server"; + info.version = "1.0"; + mcp::Server server(info, mcp::ServerCapabilities{}); + + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 2) { + raw_transport->close(); + } + }); + + raw_transport->enqueue_message(make_initialize_request("init").dump()); + auto initialized = make_initialized_notification(); + initialized["params"] = "not-an-object"; + raw_transport->enqueue_message(initialized.dump()); + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", "list"}, {"method", "tools/list"}}.dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("error")); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_INVALID_REQUEST); +} + +TEST_F(ServerCoreTest, FailedSessionTeardownDoesNotStrandTheServer) { + mcp::Server server({"teardown-server", "1.0"}, mcp::ServerCapabilities{}); + + std::exception_ptr first_error; + boost::asio::co_spawn( + io_ctx_, server.run(std::make_shared(), io_ctx_.get_executor()), + [&first_error](std::exception_ptr error) { first_error = std::move(error); }); + io_ctx_.run(); + EXPECT_NE(first_error, nullptr); + + io_ctx_.restart(); + + // The failed teardown must not leave the session registered, or every later run() is refused. + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + raw_transport->close(); + }); + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", "ping"}, {"method", "ping"}}.dump()); + + std::exception_ptr second_error; + boost::asio::co_spawn( + io_ctx_, server.run(transport, io_ctx_.get_executor()), + [&second_error](std::exception_ptr error) { second_error = std::move(error); }); + io_ctx_.run(); + + EXPECT_EQ(second_error, nullptr); + ASSERT_EQ(responses.size(), 1); + EXPECT_EQ(responses[0]["id"], "ping"); +} + +TEST_F(ServerCoreTest, InvalidJsonRpcEnvelopesAreRejected) { + mcp::Implementation info; + info.name = "test-server"; + info.version = "1.0"; + mcp::Server server(info, mcp::ServerCapabilities{}); + + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 4) { + raw_transport->close(); + } + }); + + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "1.0"}, {"id", "response"}, {"result", nlohmann::json::object()}} + .dump()); + raw_transport->enqueue_message(nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "response"}, + {"result", nlohmann::json::object()}, + {"error", {{"code", -1}, {"message", "bad"}}}} + .dump()); + raw_transport->enqueue_message( + nlohmann::json{{"id", "missing-version"}, {"method", "ping"}}.dump()); + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "1.0"}, {"id", "wrong-version"}, {"method", "ping"}}.dump()); + raw_transport->enqueue_message(nlohmann::json{{"jsonrpc", "2.0"}, {"id", "missing-method"}}.dump()); + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", "bad-method"}, {"method", 7}}.dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 4); + for (const auto& response : responses) { + EXPECT_EQ(response["error"]["code"], mcp::g_INVALID_REQUEST); + } +} + +TEST_F(ServerCoreTest, MalformedAndNonObjectMessagesReturnErrorsThenAcceptValidRequest) { + mcp::Server server({"validation-server", "1.0"}, mcp::ServerCapabilities{}); + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 3) { + raw_transport->close(); + } + }); + + raw_transport->enqueue_message(R"({"jsonrpc":"2.0","id":)"); + raw_transport->enqueue_message(nlohmann::json::array().dump()); + raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, {"id", "ping"}, {"method", "ping"}}.dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 3); + EXPECT_EQ(responses[0]["error"]["code"], mcp::g_PARSE_ERROR); + EXPECT_TRUE(responses[0]["id"].is_null()); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_INVALID_REQUEST); + EXPECT_TRUE(responses[1]["id"].is_null()); + EXPECT_EQ(responses[2]["id"], "ping"); + EXPECT_EQ(responses[2]["result"], nlohmann::json::object()); +} + +TEST_F(ServerCoreTest, MemoryTransportRecoversFromMalformedJson) { + mcp::Server server({"validation-server", "1.0"}, mcp::ServerCapabilities{}); + auto [server_transport, client_transport] = + mcp::create_memory_transport_pair(io_ctx_.get_executor()); + + std::vector responses; + std::exception_ptr server_error; + std::exception_ptr client_error; + boost::asio::co_spawn( + io_ctx_, server.run(server_transport, io_ctx_.get_executor()), + [&server_error](std::exception_ptr error) { server_error = std::move(error); }); + boost::asio::co_spawn( + io_ctx_, verify_memory_transport_parse_recovery(client_transport, &responses), + [&client_error](std::exception_ptr error) { client_error = std::move(error); }); + + io_ctx_.run(); + + EXPECT_EQ(server_error, nullptr); + EXPECT_EQ(client_error, nullptr); + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[0]["error"]["code"], mcp::g_PARSE_ERROR); + EXPECT_TRUE(responses[0]["id"].is_null()); + EXPECT_EQ(responses[1]["id"], 7); + EXPECT_EQ(responses[1]["result"], nlohmann::json::object()); +} + +TEST_F(ServerCoreTest, ParameterDecodingErrorsRemainDistinctFromHandlerFailures) { + mcp::Server server({"validation-server", "1.0"}, mcp::ServerCapabilities{}); + mcp::Resource resource; + resource.uri = "file:///failure.txt"; + resource.name = "failure"; + server.add_resource( + resource, [](mcp::ReadResourceRequestParams) -> mcp::ReadResourceResult { + throw std::runtime_error("resource handler failed"); + }); + server.set_page_size(1); + + std::vector requests = { + {{"jsonrpc", "2.0"}, + {"id", "resource-params"}, + {"method", "resources/read"}, + {"params", nlohmann::json::object()}}, + {{"jsonrpc", "2.0"}, + {"id", "subscribe-params"}, + {"method", "resources/subscribe"}, + {"params", nlohmann::json::object()}}, + {{"jsonrpc", "2.0"}, + {"id", "prompt-params"}, + {"method", "prompts/get"}, + {"params", {{"name", 7}}}}, + {{"jsonrpc", "2.0"}, + {"id", "logging-params"}, + {"method", "logging/setLevel"}, + {"params", nlohmann::json::object()}}, + {{"jsonrpc", "2.0"}, + {"id", "completion-params"}, + {"method", "completion/complete"}, + {"params", nlohmann::json::object()}}, + {{"jsonrpc", "2.0"}, + {"id", "pagination-params"}, + {"method", "resources/list"}, + {"params", {{"cursor", 7}}}}, + {{"jsonrpc", "2.0"}, + {"id", "handler"}, + {"method", "resources/read"}, + {"params", {{"uri", "file:///failure.txt"}}}}, + }; + + std::vector responses; + std::exception_ptr dispatch_error; + boost::asio::co_spawn( + io_ctx_, + [&server, &requests, &responses]() -> mcp::Task { + for (auto& request : requests) { + responses.push_back( + nlohmann::json::parse(co_await server.dispatch_request_direct(std::move(request)))); + } + }, + [&dispatch_error](std::exception_ptr error) { dispatch_error = std::move(error); }); + + io_ctx_.run(); + + EXPECT_EQ(dispatch_error, nullptr); + ASSERT_EQ(responses.size(), requests.size()); + for (std::size_t index = 0; index + 1 < responses.size(); ++index) { + EXPECT_EQ(responses[index]["error"]["code"], mcp::g_INVALID_PARAMS); + } + EXPECT_EQ(responses.back()["error"]["code"], mcp::g_INTERNAL_ERROR); + EXPECT_EQ(responses.back()["error"]["message"], "resource handler failed"); +} + +TEST_F(ServerCoreTest, DirectDispatchDoesNotRequireStatefulHandshake) { + mcp::Implementation info; + info.name = "stateless-server"; + info.version = "1.0"; + mcp::Server server(info, mcp::ServerCapabilities{}); + + nlohmann::json response; + std::exception_ptr error; + boost::asio::co_spawn(io_ctx_, + server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "list"}, {"method", "tools/list"}}), + [&response, &error](std::exception_ptr dispatch_error, std::string wire) { + error = std::move(dispatch_error); + if (error) { + return; + } + response = nlohmann::json::parse(wire); + }); + + io_ctx_.run(); + + EXPECT_EQ(error, nullptr); + EXPECT_TRUE(response["result"]["tools"].empty()); +} + +TEST_F(ServerCoreTest, InitializeRequiresCompleteTypedParameters) { + mcp::Server server({"validation-server", "1.0"}, mcp::ServerCapabilities{}); + + nlohmann::json response; + std::exception_ptr error; + boost::asio::co_spawn( + io_ctx_, + server.dispatch_request_direct(nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "init"}, + {"method", "initialize"}, + {"params", nlohmann::json::object()}}), + [&response, &error](std::exception_ptr dispatch_error, std::string wire) { + error = std::move(dispatch_error); + if (!error) { + response = nlohmann::json::parse(wire); + } + }); + + io_ctx_.run(); + + EXPECT_EQ(error, nullptr); + EXPECT_EQ(response["id"], "init"); + EXPECT_EQ(response["error"]["code"], mcp::g_INVALID_PARAMS); +} + +TEST_F(ServerCoreTest, ToolArgumentsMustBeAnObject) { + mcp::Server server({"validation-server", "1.0"}, mcp::ServerCapabilities{}); + + nlohmann::json response; + std::exception_ptr error; + boost::asio::co_spawn(io_ctx_, + server.dispatch_request_direct( + nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "tool"}, + {"method", "tools/call"}, + {"params", {{"name", "echo"}, {"arguments", nullptr}}}}), + [&response, &error](std::exception_ptr dispatch_error, std::string wire) { + error = std::move(dispatch_error); + if (!error) { + response = nlohmann::json::parse(wire); + } + }); + + io_ctx_.run(); + + EXPECT_EQ(error, nullptr); + EXPECT_EQ(response["id"], "tool"); + EXPECT_EQ(response["error"]["code"], mcp::g_INVALID_PARAMS); +} + +TEST_F(ServerCoreTest, RunWaitsForDispatchedHandlersBeforeReturning) { + using namespace std::chrono_literals; + + mcp::ServerCapabilities capabilities; + capabilities.tools = mcp::ServerCapabilities::ToolsCapability{}; + mcp::Server server({"draining-server", "1.0"}, capabilities); + auto transport = std::make_shared(io_ctx_.get_executor()); + + bool handler_started = false; + bool handler_completed = false; + server.add_tool( + "delayed", "Delayed result", nlohmann::json{{"type", "object"}}, + [this, transport, &handler_started, + &handler_completed](const nlohmann::json&) -> mcp::Task { + handler_started = true; + transport->close(); + boost::asio::steady_timer delay(io_ctx_, 10ms); + co_await delay.async_wait(boost::asio::use_awaitable); + handler_completed = true; + co_return nlohmann::json{{"done", true}}; + }); + + transport->enqueue_message(make_initialize_request("init").dump()); + transport->enqueue_message(make_initialized_notification().dump()); + transport->enqueue_message(make_tool_call_request("tool", "delayed").dump()); + + bool run_completed = false; + boost::asio::co_spawn(io_ctx_, server.run(transport, io_ctx_.get_executor()), + [&run_completed](std::exception_ptr error) { + EXPECT_EQ(error, nullptr); + run_completed = true; + }); + + io_ctx_.run(); + + EXPECT_TRUE(handler_started); + EXPECT_TRUE(handler_completed); + EXPECT_TRUE(run_completed); +} + +TEST_F(ServerCoreTest, SessionTeardownStaysOnConfiguredStrandAcrossRepeatedRuns) { + constexpr std::size_t run_count = 64; + auto configured_strand = boost::asio::make_strand(io_ctx_.get_executor()); + mcp::Server server({"threaded-server", "1.0"}, mcp::ServerCapabilities{}); + + std::vector> transports; + transports.reserve(run_count); + std::exception_ptr run_error; + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + for (std::size_t index = 0; index < run_count; ++index) { + auto transport = std::make_shared(configured_strand); + transports.push_back(transport); + co_await server.run(transport, configured_strand); + } + }, + [&run_error](std::exception_ptr error) { run_error = std::move(error); }); + + run_server_io_on_thread_pool(io_ctx_); + + EXPECT_EQ(run_error, nullptr); + ASSERT_EQ(transports.size(), run_count); + for (const auto& transport : transports) { + EXPECT_GE(transport->close_calls(), 1); + EXPECT_EQ(transport->off_strand_closes(), 0); + } +} + +// A peer that serializes absent optionals as explicit nulls answers a server-initiated request +// with an "error" member set to null. That is not an error, and the response must still correlate. +TEST_F(ServerCoreTest, ReverseRequestResponseCarryingNullErrorIsAccepted) { + using namespace std::chrono_literals; + + auto [server_transport, client_transport] = + mcp::create_memory_transport_pair(io_ctx_.get_executor()); + mcp::Server server({"reverse-rpc-server", "1.0"}, mcp::ServerCapabilities{}); + + std::atomic_bool session_ready{false}; + boost::asio::co_spawn(io_ctx_, server.run(server_transport, io_ctx_.get_executor()), + boost::asio::detached); + + boost::asio::co_spawn( + io_ctx_, + [client_transport, &session_ready]() -> mcp::Task { + co_await client_transport->write_message(make_initialize_request("init").dump()); + static_cast(co_await client_transport->read_message()); + co_await client_transport->write_message(make_initialized_notification().dump()); + session_ready.store(true, std::memory_order_release); + + auto request = nlohmann::json::parse(co_await client_transport->read_message()); + // [gcc11-sso: scope-before-await] Build the message first: a braced-init-list temporary + // inside the co_await argument makes GCC fail with an internal compiler error. + auto response = + make_result_response(request.at("id").get(), nlohmann::json{{"ok", true}}); + response["error"] = nullptr; + co_await client_transport->write_message(response.dump()); + }, + boost::asio::detached); + + bool completed = false; + nlohmann::json reverse_result; + boost::asio::steady_timer launch_poll(io_ctx_); + std::function launch; + launch = [&]() { + if (!session_ready.load(std::memory_order_acquire)) { + launch_poll.expires_after(1ms); + launch_poll.async_wait([&launch](const boost::system::error_code& error) { + if (!error) { + launch(); + } + }); + return; + } + boost::asio::co_spawn( + io_ctx_, server.send_request("sampling/createMessage", nlohmann::json::object()), + [this, &completed, &reverse_result](std::exception_ptr error, nlohmann::json result) { + if (!error) { + completed = true; + reverse_result = std::move(result); + } + // The live session keeps the context busy, so stop it here; otherwise run_for + // waits out the whole timeout even when the response arrived immediately. + io_ctx_.stop(); + }); + }; + launch(); + + io_ctx_.run_for(5s); + + ASSERT_TRUE(completed) << "the reverse request never correlated with its response"; + EXPECT_EQ(reverse_result["ok"], true); +} + +TEST_F(ServerCoreTest, ConcurrentReverseRequestsRemainCorrelatedOnMultiThreadedExecutor) { + using namespace std::chrono_literals; + + constexpr std::size_t request_count = 256; + auto [server_transport, client_transport] = + mcp::create_memory_transport_pair(io_ctx_.get_executor()); + mcp::Server server({"reverse-rpc-server", "1.0"}, mcp::ServerCapabilities{}); + + std::atomic_bool session_ready{false}; + std::atomic_bool stop_poll{false}; + std::atomic_size_t completed{0}; + std::atomic_size_t correct{0}; + std::atomic_size_t errors{0}; + + boost::asio::co_spawn(io_ctx_, server.run(server_transport, io_ctx_.get_executor()), + [&errors](std::exception_ptr error) { + if (error) { + errors.fetch_add(1, std::memory_order_relaxed); + } + }); + + boost::asio::co_spawn( + io_ctx_, + [client_transport, &session_ready]() -> mcp::Task { + co_await client_transport->write_message(make_initialize_request("init").dump()); + static_cast(co_await client_transport->read_message()); + co_await client_transport->write_message(make_initialized_notification().dump()); + session_ready.store(true, std::memory_order_release); + + for (std::size_t response_count = 0; response_count < request_count;) { + auto request = nlohmann::json::parse(co_await client_transport->read_message()); + if (!request.contains("id")) { + continue; + } + auto result = nlohmann::json{{"sequence", request.at("params").at("sequence")}}; + co_await client_transport->write_message( + make_result_response(request.at("id").get(), std::move(result)) + .dump()); + ++response_count; + } + }, + [&errors](std::exception_ptr error) { + if (error) { + errors.fetch_add(1, std::memory_order_relaxed); + } + }); + + boost::asio::steady_timer launch_poll(io_ctx_); + std::function launch_requests; + launch_requests = [&]() { + if (!session_ready.load(std::memory_order_acquire)) { + launch_poll.expires_after(1ms); + launch_poll.async_wait([&launch_requests](const boost::system::error_code& error) { + if (!error) { + launch_requests(); + } + }); + return; + } + + for (std::size_t sequence = 0; sequence < request_count; ++sequence) { + boost::asio::co_spawn( + io_ctx_, + server.send_request("sampling/createMessage", nlohmann::json{{"sequence", sequence}}), + [sequence, &completed, &correct, &errors](std::exception_ptr error, + nlohmann::json result) { + if (error) { + errors.fetch_add(1, std::memory_order_relaxed); + } else if (result.at("sequence").get() == sequence) { + correct.fetch_add(1, std::memory_order_relaxed); + } + completed.fetch_add(1, std::memory_order_release); + }); + } + }; + launch_requests(); + + // Wide for the same reason as in the client's twin of this test: a sanitizer can stop every + // thread for seconds, and the bound only ends a run that has gone wrong. + boost::asio::steady_timer watchdog(io_ctx_); + watchdog.expires_after(60s); + watchdog.async_wait( + [&completed, &stop_poll, client_transport](const boost::system::error_code& error) { + if (!error && completed.load(std::memory_order_acquire) != request_count) { + stop_poll.store(true, std::memory_order_release); + client_transport->close(); + } + }); + + boost::asio::steady_timer completion_poll(io_ctx_); + std::function poll_completion; + poll_completion = [&]() { + if (stop_poll.load(std::memory_order_acquire)) { + return; + } + if (completed.load(std::memory_order_acquire) == request_count) { + stop_poll.store(true, std::memory_order_release); + watchdog.cancel(); + client_transport->close(); + return; + } + completion_poll.expires_after(1ms); + completion_poll.async_wait([&poll_completion](const boost::system::error_code& error) { + if (!error) { + poll_completion(); + } + }); + }; + poll_completion(); + + run_server_io_on_thread_pool(io_ctx_); + + EXPECT_EQ(completed.load(), request_count); + EXPECT_EQ(correct.load(), request_count); + EXPECT_EQ(errors.load(), 0); +} + +// --------------------------------------------------------------------------- +// server/discover +// --------------------------------------------------------------------------- + +TEST_F(ServerCoreTest, DiscoverWithoutInitializeReturnsFullPayload) { + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + + mcp::Implementation server_info; + server_info.name = "test-server"; + server_info.version = "1.0"; + + mcp::ServerCapabilities caps; + mcp::ServerCapabilities::ToolsCapability tools_cap; + tools_cap.listChanged = true; + caps.tools = std::move(tools_cap); + + mcp::Server server(std::move(server_info), std::move(caps)); + + nlohmann::json response; + raw_transport->set_on_write([&response, raw_transport](std::string_view msg) { + response = nlohmann::json::parse(msg); + raw_transport->close(); + }); + + nlohmann::json discover_req = {{"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}; + raw_transport->enqueue_message(discover_req.dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_TRUE(response.contains("result")); + EXPECT_EQ(response["id"], "d1"); + + auto result = response["result"]; + EXPECT_EQ(result["resultType"], "complete"); + + std::vector expected_versions = {"2024-11-05", "2025-03-26", "2025-06-18", + "2025-11-25", "2026-07-28"}; + EXPECT_EQ(result["supportedVersions"].get>(), expected_versions); + + ASSERT_TRUE(result["capabilities"].contains("tools")); + EXPECT_TRUE(result["capabilities"]["tools"]["listChanged"]); + + ASSERT_TRUE(result.contains("_meta")); + ASSERT_TRUE(result["_meta"].contains("io.modelcontextprotocol/serverInfo")); + EXPECT_EQ(result["_meta"]["io.modelcontextprotocol/serverInfo"]["name"], "test-server"); + EXPECT_EQ(result["_meta"]["io.modelcontextprotocol/serverInfo"]["version"], "1.0"); + + // Not initialized: server/discover is a pre-gate method and must not require or mutate + // lifecycle state, reachable with zero prior state. + EXPECT_FALSE(server.is_initialized()); +} + +TEST_F(ServerCoreTest, DiscoverThenInitializeStillNegotiatesLegacyVersionsUnchanged) { + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + + mcp::Implementation server_info; + server_info.name = "test-server"; + server_info.version = "1.0"; + + mcp::Server server(std::move(server_info), mcp::ServerCapabilities{}); + + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view msg) { + responses.push_back(nlohmann::json::parse(msg)); + if (responses.size() == 2) { + raw_transport->close(); + } + }); + + nlohmann::json discover_req = {{"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}; + nlohmann::json init_req = make_initialize_request("i1"); + // An unknown/unsupported protocol version must still negotiate to the latest legacy + // version, byte-identical to pre-discover behavior. + init_req["params"]["protocolVersion"] = "1900-01-01"; + + raw_transport->enqueue_message(discover_req.dump()); + raw_transport->enqueue_message(init_req.dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[0].contains("result")); + EXPECT_EQ(responses[0]["result"]["resultType"], "complete"); + + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_EQ(responses[1]["result"]["protocolVersion"], "2025-11-25"); + EXPECT_EQ(responses[1]["result"]["protocolVersion"], std::string(mcp::g_LATEST_PROTOCOL_VERSION)); + + EXPECT_TRUE(server.is_initialized()); +} + +TEST_F(ServerCoreTest, DiscoverAcceptsMetaWithoutMutatingLifecycle) { + mcp::Server server({"meta-server", "1.0"}, mcp::ServerCapabilities{}); + + nlohmann::json response; + std::exception_ptr error; + nlohmann::json discover_req = { + {"jsonrpc", "2.0"}, + {"id", "d1"}, + {"method", "server/discover"}, + {"params", + {{"_meta", + {{"io.modelcontextprotocol/protocolVersion", "2026-07-28"}, + {"io.modelcontextprotocol/clientInfo", {{"name", "probe-client"}, {"version", "0.1"}}}, + {"io.modelcontextprotocol/clientCapabilities", nlohmann::json::object()}}}}}}; + + boost::asio::co_spawn(io_ctx_, server.dispatch_request_direct(discover_req), + [&response, &error](std::exception_ptr dispatch_error, std::string wire) { + error = std::move(dispatch_error); + if (!error) { + response = nlohmann::json::parse(wire); + } + }); + + io_ctx_.run(); + + EXPECT_EQ(error, nullptr); + ASSERT_TRUE(response.contains("result")); + EXPECT_EQ(response["result"]["resultType"], "complete"); + + // params carrying _meta must be accepted, not rejected, and must not alter lifecycle state. + EXPECT_FALSE(server.is_initialized()); +} + +TEST_F(ServerCoreTest, DiscoverAcceptsAbsentOrEmptyParamsButRejectsNonObjectParams) { + mcp::Server server({"params-server", "1.0"}, mcp::ServerCapabilities{}); + + // params absent entirely. + nlohmann::json response_no_params; + boost::asio::co_spawn(io_ctx_, + server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}), + [&response_no_params](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + response_no_params = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + ASSERT_TRUE(response_no_params.contains("result")); + EXPECT_EQ(response_no_params["result"]["resultType"], "complete"); + + io_ctx_.restart(); + + // params present but an empty object. + nlohmann::json response_empty_params; + boost::asio::co_spawn( + io_ctx_, + server.dispatch_request_direct(nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "d2"}, + {"method", "server/discover"}, + {"params", nlohmann::json::object()}}), + [&response_empty_params](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + response_empty_params = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + ASSERT_TRUE(response_empty_params.contains("result")); + EXPECT_EQ(response_empty_params["result"]["resultType"], "complete"); + + io_ctx_.restart(); + + // params present but not an object: rejected pre-dispatch as an invalid request, the same + // as for any other method (see validate_request_envelope). + nlohmann::json response_bad_params; + boost::asio::co_spawn(io_ctx_, + server.dispatch_request_direct(nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "d3"}, + {"method", "server/discover"}, + {"params", "not-an-object"}}), + [&response_bad_params](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + response_bad_params = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + ASSERT_TRUE(response_bad_params.contains("error")); + EXPECT_EQ(response_bad_params["error"]["code"], mcp::g_INVALID_REQUEST); +} + +// server/utilities/caching.md requires servers to include caching hints on every "complete" +// result, server/discover listed first; an absent ttlMs "should only occur in older server +// versions". This SDK therefore always emits ttlMs (default 0) and cacheScope (default +// "private", the conservative choice absent a spec-stated default) unless overridden. +TEST_F(ServerCoreTest, DiscoverCachingHintsPresentWithDefaultsOverriddenWhenConfigured) { + mcp::Server default_server({"default-server", "1.0"}, mcp::ServerCapabilities{}); + + nlohmann::json default_response; + boost::asio::co_spawn(io_ctx_, + default_server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}), + [&default_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + default_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + + ASSERT_TRUE(default_response.contains("result")); + ASSERT_TRUE(default_response["result"].contains("ttlMs")); + EXPECT_EQ(default_response["result"]["ttlMs"], 0); + ASSERT_TRUE(default_response["result"].contains("cacheScope")); + EXPECT_EQ(default_response["result"]["cacheScope"], "private"); + + io_ctx_.restart(); + + mcp::Server configured_server({"configured-server", "1.0"}, mcp::ServerCapabilities{}); + configured_server.set_discover_ttl_ms(3600000); + configured_server.set_discover_cache_scope(mcp::CacheScope::ePublic); + + nlohmann::json configured_response; + boost::asio::co_spawn(io_ctx_, + configured_server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "d2"}, {"method", "server/discover"}}), + [&configured_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + configured_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + + ASSERT_TRUE(configured_response.contains("result")); + ASSERT_TRUE(configured_response["result"].contains("ttlMs")); + EXPECT_EQ(configured_response["result"]["ttlMs"], 3600000); + ASSERT_TRUE(configured_response["result"].contains("cacheScope")); + EXPECT_EQ(configured_response["result"]["cacheScope"], "public"); + + io_ctx_.restart(); + + // Explicitly configuring CacheScope::ePrivate must also serialize correctly, distinct from + // it merely being the unconfigured default asserted above. + mcp::Server explicit_private_server({"explicit-private-server", "1.0"}, mcp::ServerCapabilities{}); + explicit_private_server.set_discover_cache_scope(mcp::CacheScope::ePrivate); + + nlohmann::json explicit_private_response; + boost::asio::co_spawn(io_ctx_, + explicit_private_server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "d3"}, {"method", "server/discover"}}), + [&explicit_private_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + explicit_private_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + + ASSERT_TRUE(explicit_private_response.contains("result")); + ASSERT_TRUE(explicit_private_response["result"].contains("cacheScope")); + EXPECT_EQ(explicit_private_response["result"]["cacheScope"], "private"); +} + +TEST_F(ServerCoreTest, SetDiscoverTtlMsRejectsNegativeValues) { + mcp::Server server({"negative-ttl-server", "1.0"}, mcp::ServerCapabilities{}); + EXPECT_THROW(server.set_discover_ttl_ms(-1), std::invalid_argument); +} + +TEST_F(ServerCoreTest, InstructionsAppearInBothInitializeAndDiscoverWhenSet) { + mcp::Server server({"instructed-server", "1.0"}, mcp::ServerCapabilities{}); + server.set_instructions("Use the tools wisely."); + + nlohmann::json discover_response; + boost::asio::co_spawn(io_ctx_, + server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}), + [&discover_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + discover_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + ASSERT_TRUE(discover_response.contains("result")); + ASSERT_TRUE(discover_response["result"].contains("instructions")); + EXPECT_EQ(discover_response["result"]["instructions"], "Use the tools wisely."); + + io_ctx_.restart(); + + nlohmann::json init_response; + boost::asio::co_spawn(io_ctx_, server.dispatch_request_direct(make_initialize_request("i1")), + [&init_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + init_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + ASSERT_TRUE(init_response.contains("result")); + ASSERT_TRUE(init_response["result"].contains("instructions")); + EXPECT_EQ(init_response["result"]["instructions"], "Use the tools wisely."); +} + +TEST_F(ServerCoreTest, InstructionsOmittedFromBothInitializeAndDiscoverWhenUnset) { + mcp::Server server({"uninstructed-server", "1.0"}, mcp::ServerCapabilities{}); + + nlohmann::json discover_response; + boost::asio::co_spawn(io_ctx_, + server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}), + [&discover_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + discover_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + ASSERT_TRUE(discover_response.contains("result")); + EXPECT_FALSE(discover_response["result"].contains("instructions")); + + io_ctx_.restart(); + + nlohmann::json init_response; + boost::asio::co_spawn(io_ctx_, server.dispatch_request_direct(make_initialize_request("i1")), + [&init_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + init_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + ASSERT_TRUE(init_response.contains("result")); + EXPECT_FALSE(init_response["result"].contains("instructions")); +} + +TEST_F(ServerCoreTest, DiscoverCapabilitiesMatchInitializeCapabilitiesForNonTrivialSet) { + mcp::ServerCapabilities caps; + mcp::ServerCapabilities::ToolsCapability tools_cap; + tools_cap.listChanged = true; + caps.tools = tools_cap; + mcp::ServerCapabilities::ResourcesCapability resources_cap; + resources_cap.listChanged = true; + resources_cap.subscribe = false; + caps.resources = resources_cap; + mcp::ServerCapabilities::PromptsCapability prompts_cap; + prompts_cap.listChanged = true; + caps.prompts = prompts_cap; + caps.logging = nlohmann::json::object(); + + mcp::Server server({"capable-server", "1.0"}, caps); + + nlohmann::json discover_response; + boost::asio::co_spawn(io_ctx_, + server.dispatch_request_direct(nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}), + [&discover_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + discover_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + + io_ctx_.restart(); + + nlohmann::json init_response; + boost::asio::co_spawn(io_ctx_, server.dispatch_request_direct(make_initialize_request("i1")), + [&init_response](std::exception_ptr error, std::string wire) { + EXPECT_EQ(error, nullptr); + init_response = nlohmann::json::parse(wire); + }); + io_ctx_.run(); + + ASSERT_TRUE(discover_response.contains("result")); + ASSERT_TRUE(init_response.contains("result")); + + const auto& discover_caps = discover_response["result"]["capabilities"]; + const auto& init_caps = init_response["result"]["capabilities"]; + ASSERT_TRUE(discover_caps.contains("tools")); + ASSERT_TRUE(discover_caps.contains("resources")); + ASSERT_TRUE(discover_caps.contains("prompts")); + ASSERT_TRUE(discover_caps.contains("logging")); + EXPECT_EQ(discover_caps, init_caps); +} + +TEST_F(ServerCoreTest, DiscoverWorksAgainAfterInitializeIdempotently) { + auto transport = std::make_shared(io_ctx_.get_executor()); + auto* raw_transport = transport.get(); + + mcp::Implementation server_info; + server_info.name = "test-server"; + server_info.version = "1.0"; + + mcp::Server server(std::move(server_info), mcp::ServerCapabilities{}); + + std::vector responses; + raw_transport->set_on_write([&responses, raw_transport](std::string_view msg) { + responses.push_back(nlohmann::json::parse(msg)); + if (responses.size() == 4) { + raw_transport->close(); + } + }); + + nlohmann::json discover_req_1 = {{"jsonrpc", "2.0"}, {"id", "d1"}, {"method", "server/discover"}}; + nlohmann::json discover_req_2 = {{"jsonrpc", "2.0"}, {"id", "d2"}, {"method", "server/discover"}}; + nlohmann::json ping_req = {{"jsonrpc", "2.0"}, {"id", "p1"}, {"method", "ping"}}; + + raw_transport->enqueue_message(make_initialize_request("i1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); + raw_transport->enqueue_message(discover_req_1.dump()); + raw_transport->enqueue_message(discover_req_2.dump()); + raw_transport->enqueue_message(ping_req.dump()); + + boost::asio::co_spawn( + io_ctx_, [&]() -> mcp::Task { co_await server.run(transport, io_ctx_.get_executor()); }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 4); + EXPECT_EQ(responses[0]["id"], "i1"); + ASSERT_TRUE(responses[0].contains("result")); + + EXPECT_EQ(responses[1]["id"], "d1"); + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_EQ(responses[1]["result"]["resultType"], "complete"); + + // Idempotent: a second discover call returns the exact same result (modulo the envelope id). + EXPECT_EQ(responses[2]["id"], "d2"); + ASSERT_TRUE(responses[2].contains("result")); + EXPECT_EQ(responses[2]["result"], responses[1]["result"]); + + // A follow-up request after discover still requires (and gets) the eReady lifecycle state + // reached by initialize+initialized: discover neither reset nor otherwise mutated it. + EXPECT_EQ(responses[3]["id"], "p1"); + ASSERT_TRUE(responses[3].contains("result")); + + EXPECT_TRUE(server.is_initialized()); +} + +namespace { + +// Drives one server-initiated request that the client answers with a JSON-RPC error object, and +// reports the diagnostic the server raised. Here the untrusted side is the client: every +// server-initiated request -- sampling/createMessage, elicitation, roots/list -- can be answered with +// an error whose `message` is deserialized verbatim and interpolated into the runtime_error the +// server operator logs. +struct ReverseErrorOutcome { + bool threw{false}; + std::string what; +}; + +// A free coroutine rather than a lambda, and every operand named rather than a temporary: GCC 13 +// ICEs (build_special_member_call) on a co_await of a call taking a temporary inside a try block +// in a capturing coroutine lambda. +mcp::Task capture_reverse_request_error(mcp::Server& server, ReverseErrorOutcome& outcome) { + const std::optional params{nlohmann::json{{"probe", true}}}; + try { + auto result = co_await server.send_request("sampling/createMessage", params); + static_cast(result); + } catch (const std::exception& error) { + outcome.threw = true; + outcome.what = error.what(); + } +} + +ReverseErrorOutcome reverse_request_against_peer_error(boost::asio::io_context& io_ctx, + const std::string& peer_message) { + using namespace std::chrono_literals; + + // Named copies rather than a structured binding: GCC 13 ICEs on capturing a structured binding + // by copy in a coroutine lambda alongside a by-reference default. + auto transport_pair = mcp::create_memory_transport_pair(io_ctx.get_executor()); + auto server_transport = transport_pair.first; + auto client_transport = transport_pair.second; + mcp::Server server({"reverse-error-server", "1.0"}, mcp::ServerCapabilities{}); + + ReverseErrorOutcome outcome; + std::atomic_bool session_ready{false}; + + boost::asio::co_spawn(io_ctx, server.run(server_transport, io_ctx.get_executor()), + boost::asio::detached); + + boost::asio::co_spawn( + io_ctx, + [client_transport, peer_message, &session_ready]() -> mcp::Task { + co_await client_transport->write_message(make_initialize_request("init").dump()); + static_cast(co_await client_transport->read_message()); + co_await client_transport->write_message(make_initialized_notification().dump()); + session_ready.store(true, std::memory_order_release); + + // Answer the server's reverse request with an error of the client's own choosing. The + // message travels as JSON, so CR/LF written here arrive as real control bytes: the + // encoder escapes them and the server's decoder turns them back. + for (;;) { + const auto request = nlohmann::json::parse(co_await client_transport->read_message()); + if (!request.contains("id")) { + continue; + } + co_await client_transport->write_message( + make_error_response(request.at("id").get(), mcp::g_INTERNAL_ERROR, + peer_message) + .dump()); + co_return; + } + }, + boost::asio::detached); + + boost::asio::steady_timer watchdog(io_ctx); + watchdog.expires_after(10s); + watchdog.async_wait([client_transport](const boost::system::error_code& error) { + if (!error) { + client_transport->close(); + } + }); + + boost::asio::co_spawn( + io_ctx, + [&io_ctx, &server, &outcome, &session_ready, &watchdog, client_transport]() -> mcp::Task { + boost::asio::steady_timer poll(io_ctx); + for (int attempt = 0; attempt < 5000 && !session_ready.load(std::memory_order_acquire); + ++attempt) { + poll.expires_after(1ms); + co_await poll.async_wait(boost::asio::use_awaitable); + } + co_await capture_reverse_request_error(server, outcome); + // Cancelled rather than left pending: an armed watchdog is outstanding work, and + // io_context::run() would sit on it for the full ten seconds after the test is done. + watchdog.cancel(); + client_transport->close(); + }, + boost::asio::detached); + + io_ctx.run(); + return outcome; +} + +} // namespace + +// The message in the error a CLIENT returns for a server-initiated request is the client's text, +// and it reaches the diagnostic the server operator logs. A client answering with CR/LF forges a +// line in that log; a bidi override reorders the tail of the message. +TEST_F(ServerCoreTest, ReverseRequestErrorDiagnosticFlattensThePeerChosenMessage) { + // "zqtripwire" is a token no other code path produces; see the tripwire assertions below. + // \xe2\x80\xae is U+202E RIGHT-TO-LEFT OVERRIDE, written escaped so this source file does not + // itself contain a bidi override. + const std::string forged = + "denied\r\n2026-09-20 INFO zqtripwire operator approved\xe2\x80\xae reordered tail"; + + const auto outcome = reverse_request_against_peer_error(io_ctx_, forged); + + // Tripwire, ordered before the flattening assertions. send_request() has four other throw + // sites -- "reverse RPC is unavailable in stateless direct dispatch", "server session is + // closing", "duplicate pending request id", "pending request not found for id" -- none of + // which interpolates a byte the peer chose, and each of which would satisfy everything below + // for a reason unrelated to the site under test. Dump what() on failure so a vacuous pass + // cannot hide. + ASSERT_TRUE(outcome.threw) << "send_request() raised nothing at all"; + ASSERT_EQ(outcome.what.rfind("JSON-RPC error ", 0), 0U) << "actual what(): " << outcome.what; + ASSERT_NE(outcome.what.find(std::to_string(mcp::g_INTERNAL_ERROR)), std::string::npos) + << "actual what(): " << outcome.what; + ASSERT_NE(outcome.what.find("zqtripwire"), std::string::npos) << "actual what(): " << outcome.what; + + EXPECT_EQ(outcome.what.find('\r'), std::string::npos) << "actual what(): " << outcome.what; + EXPECT_EQ(outcome.what.find('\n'), std::string::npos) << "actual what(): " << outcome.what; + EXPECT_EQ(outcome.what.find("\xe2\x80\xae"), std::string::npos) + << "actual what(): " << outcome.what; +} + +// The same site must bound the value too, so a client cannot flood the server operator's log +// through a megabyte-long error message. +TEST_F(ServerCoreTest, ReverseRequestErrorDiagnosticBoundsThePeerChosenMessage) { + const std::string forged = "zqtripwire" + std::string(64 * 1024, 'A'); + + const auto outcome = reverse_request_against_peer_error(io_ctx_, forged); + + ASSERT_TRUE(outcome.threw) << "send_request() raised nothing at all"; + ASSERT_EQ(outcome.what.rfind("JSON-RPC error ", 0), 0U) + << "actual what() prefix: " << outcome.what.substr(0, 80); + ASSERT_NE(outcome.what.find("zqtripwire"), std::string::npos) + << "actual what() prefix: " << outcome.what.substr(0, 80); + + EXPECT_LT(outcome.what.size(), forged.size()) << "what() size: " << outcome.what.size(); + EXPECT_LE(outcome.what.size(), std::size_t{512}) << "what() size: " << outcome.what.size(); +} + +// --------------------------------------------------------------------------- +// Server destruction +// --------------------------------------------------------------------------- + +namespace { + +// A transport whose read loop stays suspended for as long as the test needs it to. close() is +// counted but deliberately does not wake the reader, so the session loop never resumes after the +// Server it belongs to is gone. Resuming it would exercise a different rule -- run() documents +// that the Server must outlive the task it returns -- and these tests are about what the +// destructor itself does on the thread that runs it. +class QuiescentSessionTransport final : public mcp::ITransport { + public: + explicit QuiescentSessionTransport(const boost::asio::any_io_executor& executor) + : reader_(executor) { + reader_.expires_at(std::chrono::steady_clock::time_point::max()); + } + + mcp::Task read_message() override { + read_calls_.fetch_add(1, std::memory_order_release); + for (;;) { + { + std::lock_guard lock(mutex_); + if (!incoming_.empty()) { + auto message = std::move(incoming_.front()); + incoming_.pop(); + co_return message; + } + } + try { + co_await reader_.async_wait(boost::asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != boost::asio::error::operation_aborted) { + throw; + } + } + } + } + + // Callable from any thread: the timer is only ever touched on its own executor. + void feed(std::string message) { + { + std::lock_guard lock(mutex_); + incoming_.push(std::move(message)); + } + boost::asio::post(reader_.get_executor(), [this] { reader_.cancel(); }); + } + + mcp::Task write_message(std::string_view message) override { + if (message.find("sampling/createMessage") != std::string_view::npos) { + reverse_requests_.fetch_add(1, std::memory_order_release); + } + co_return; + } + + void close() override { close_calls_.fetch_add(1, std::memory_order_release); } + + // Non-zero once run_session has registered the session and asked for its first message. + [[nodiscard]] std::size_t read_calls() const { return read_calls_.load(std::memory_order_acquire); } + [[nodiscard]] std::size_t reverse_requests() const { + return reverse_requests_.load(std::memory_order_acquire); + } + [[nodiscard]] std::size_t close_calls() const { + return close_calls_.load(std::memory_order_acquire); + } + + private: + boost::asio::steady_timer reader_; + std::mutex mutex_; + std::queue incoming_; + std::atomic_size_t read_calls_{0}; + std::atomic_size_t reverse_requests_{0}; + std::atomic_size_t close_calls_{0}; +}; + +// Runs `body` on a helper thread and aborts the process if it has not returned within `budget`. +// A destructor that blocks forever hangs the whole test binary rather than failing one test, and +// a hung binary reports neither a pass nor a failure for anything in it. Aborting converts that +// into a loud failure with a named cause. std::abort is used instead of a gtest failure because +// the blocked thread cannot be joined or unwound. +void run_with_teardown_watchdog(std::chrono::milliseconds budget, const std::function& body) { + std::promise finished; + auto reached_the_end = finished.get_future(); + std::thread worker([&finished, &body]() { + try { + body(); + } catch (...) { + finished.set_exception(std::current_exception()); + return; + } + finished.set_value(); + }); + + if (reached_the_end.wait_for(budget) != std::future_status::ready) { + std::fprintf(stderr, "server teardown did not finish within %lld ms\n", + static_cast(budget.count())); + std::fflush(stderr); + std::abort(); + } + + worker.join(); + reached_the_end.get(); +} + +// Spins until `condition` holds, or the budget runs out. Returns whether it held. +bool wait_for_condition(const std::function& condition, std::chrono::milliseconds budget) { + const auto deadline = std::chrono::steady_clock::now() + budget; + while (!condition()) { + if (std::chrono::steady_clock::now() >= deadline) { + return false; + } + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + return true; +} + +constexpr std::chrono::milliseconds g_teardown_budget{10000}; + +} // namespace + +// Positive control for the watchdog the tests below rely on: a body that deliberately overruns +// its budget must abort with a named cause. Without this, a watchdog that had stopped biting +// would make every test below look like it passed. Disabled because it aborts on purpose; run it +// with --gtest_also_run_disabled_tests to re-check the watchdog. +TEST(ServerTeardownTest, DISABLED_TeardownWatchdogAbortsOnOverrun) { + run_with_teardown_watchdog(std::chrono::milliseconds(200), + [] { std::this_thread::sleep_for(std::chrono::seconds(30)); }); +} + +// Building a Server, registering tools and letting it go out of scope without ever running it is +// the ordinary lifecycle in this SDK's own tests and examples. The destructor therefore has to +// complete with no executor to post to and no session to unwind. Returning is the whole +// assertion: what this detects is a destructor that blocks, which the watchdog turns into an +// abort rather than a hang. +TEST(ServerTeardownTest, DestroyWithoutSessionCompletes) { + run_with_teardown_watchdog(g_teardown_budget, [] { + mcp::ServerCapabilities capabilities; + capabilities.tools = mcp::ServerCapabilities::ToolsCapability{}; + mcp::Server server({"teardown-no-session", "1.0"}, capabilities); + server.add_tool( + "noop", "Does nothing", nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> mcp::Task { + co_return nlohmann::json::object(); + }); + }); +} + +// run() was spawned but the io_context was never run, so the session was never created and the +// posted work never executed. Anything the destructor posted here would never be picked up. +// Returning is the whole assertion, as above. +TEST(ServerTeardownTest, DestroyWithNeverRunIoContextCompletes) { + run_with_teardown_watchdog(g_teardown_budget, [] { + boost::asio::io_context io_ctx; + auto transport = std::make_shared(io_ctx.get_executor()); + mcp::Server server({"teardown-never-run", "1.0"}, mcp::ServerCapabilities{}); + boost::asio::co_spawn(io_ctx, server.run(transport, io_ctx.get_executor()), + boost::asio::detached); + }); +} + +// The session exists and holds pending reverse requests, but the io_context has been stopped and +// its thread joined, so nothing can run on the session strand again. A destructor that waited for +// strand work to complete would never be satisfied here. +TEST(ServerTeardownTest, DestroyAfterIoContextStoppedCompletes) { + boost::asio::io_context io_ctx; + auto transport = std::make_shared(io_ctx.get_executor()); + auto server = std::make_unique(mcp::Implementation{"teardown-stopped", "1.0"}, + mcp::ServerCapabilities{}); + + boost::asio::co_spawn(io_ctx, server->run(transport, io_ctx.get_executor()), boost::asio::detached); + + auto work = boost::asio::make_work_guard(io_ctx); + std::thread runner([&io_ctx] { io_ctx.run(); }); + + ASSERT_TRUE(wait_for_condition([&] { return transport->read_calls() >= 1; }, g_teardown_budget)); + + std::atomic_size_t settled{0}; + boost::asio::co_spawn( + io_ctx, server->send_request("sampling/createMessage", nlohmann::json{{"sequence", 0}}), + [&settled](std::exception_ptr, nlohmann::json) { + settled.fetch_add(1, std::memory_order_release); + }); + ASSERT_TRUE( + wait_for_condition([&] { return transport->reverse_requests() == 1; }, g_teardown_budget)); + + work.reset(); + io_ctx.stop(); + runner.join(); + + ASSERT_EQ(settled.load(std::memory_order_acquire), 0U); + + run_with_teardown_watchdog(g_teardown_budget, [&server] { server.reset(); }); +} + +// Destruction from a thread that is itself driving the session's executor. With a single-threaded +// io_context the destroying thread is the thread the session strand runs on, so a destructor that +// waited on strand work would deadlock against itself. +TEST(ServerTeardownTest, DestroyFromSessionExecutorThreadCompletes) { + boost::asio::io_context io_ctx; + auto transport = std::make_shared(io_ctx.get_executor()); + auto server = std::make_unique(mcp::Implementation{"teardown-on-executor", "1.0"}, + mcp::ServerCapabilities{}); + + boost::asio::co_spawn(io_ctx, server->run(transport, io_ctx.get_executor()), boost::asio::detached); + + run_with_teardown_watchdog(g_teardown_budget, [&] { + while (transport->read_calls() == 0 && io_ctx.poll_one() > 0) { + } + server.reset(); + }); + + EXPECT_GE(transport->read_calls(), 1U) << "the session was never registered"; + EXPECT_GE(transport->close_calls(), 1U); +} + +// The scenario: a Server destroyed from a thread that is not the session strand, while other +// threads are servicing that strand and the session's pending-request map is full. +// +// The two ASSERTs before the destructor are the reachability witness, and they stand on their own +// without a sanitizer: every reverse request has reached the wire, so each owns an entry in the +// session's pending-request map, and none has settled, so none of those entries has been erased. +// The map the destructor is about to walk therefore holds pending_count entries. +TEST(ServerTeardownTest, DestroyFromForeignThreadWithPendingReverseRequests) { + constexpr std::size_t pending_count = 512; + constexpr std::size_t pool_size = 2; + + boost::asio::io_context io_ctx; + auto transport = std::make_shared(io_ctx.get_executor()); + auto server = std::make_unique(mcp::Implementation{"teardown-foreign", "1.0"}, + mcp::ServerCapabilities{}); + + boost::asio::co_spawn(io_ctx, server->run(transport, io_ctx.get_executor()), boost::asio::detached); + + auto work = boost::asio::make_work_guard(io_ctx); + std::vector pool; + pool.reserve(pool_size); + for (std::size_t index = 0; index < pool_size; ++index) { + pool.emplace_back([&io_ctx] { io_ctx.run(); }); + } + + ASSERT_TRUE(wait_for_condition([&] { return transport->read_calls() >= 1; }, g_teardown_budget)); + + std::atomic_size_t settled{0}; + for (std::size_t sequence = 0; sequence < pending_count; ++sequence) { + boost::asio::co_spawn( + io_ctx, + server->send_request("sampling/createMessage", nlohmann::json{{"sequence", sequence}}), + [&settled](std::exception_ptr, nlohmann::json) { + settled.fetch_add(1, std::memory_order_release); + }); + } + + ASSERT_TRUE(wait_for_condition([&] { return transport->reverse_requests() == pending_count; }, + g_teardown_budget)); + ASSERT_EQ(settled.load(std::memory_order_acquire), 0U); + ASSERT_EQ(transport->close_calls(), 0U); + + run_with_teardown_watchdog(g_teardown_budget, [&server] { server.reset(); }); + + // Closing the transport is the last step of session teardown, after the pending-request walk, + // so observing it proves the teardown ran over the non-empty map witnessed above. + EXPECT_GE(transport->close_calls(), 1U); + + // Every pending request is failed and woken even though run_session never reaches its own + // teardown here: this transport does not wake its reader on close, so whatever the destructor + // arranges is the only cleanup that runs. Destroying a Server mid-session must not strand the + // callers waiting on its reverse requests. + EXPECT_TRUE(wait_for_condition( + [&] { return settled.load(std::memory_order_acquire) == pending_count; }, g_teardown_budget)) + << "settled " << settled.load(std::memory_order_acquire) << " of " << pending_count; + + work.reset(); + io_ctx.stop(); + for (auto& thread : pool) { + thread.join(); + } +} + +// A handler that is still running when its Server is destroyed must be able to find that out +// instead of dereferencing freed memory. The reverse-RPC call below is the reachable hazard: it +// goes through the Context the handler was given, which is the only route application code has +// back into the Server. +// +// The failure it reports has to be distinguishable from the two neighbouring ones -- stateless +// direct dispatch, and a session that is merely closing -- because a handler's correct response +// differs in each case, so the message is asserted rather than just the fact of an exception. +TEST(ServerTeardownTest, ReverseRequestFromAHandlerReportsADestroyedServer) { + boost::asio::io_context io_ctx; + auto transport = std::make_shared(io_ctx.get_executor()); + auto server = std::make_unique(mcp::Implementation{"teardown-handler", "1.0"}, + mcp::ServerCapabilities{}); + + std::atomic_bool handler_entered{false}; + std::atomic_bool server_destroyed{false}; + std::atomic_bool call_returned{false}; + std::string observed; + + // A tool handler is the only way application code is handed a Context, so the hazard is + // reached the way a real one would reach it. The handler parks until the Server is gone, then + // calls back through the Context it was given. + server->add_tool( + "park", "Waits for the server to go away", nlohmann::json{{"type", "object"}}, + [&](mcp::Context& ctx, const nlohmann::json&) -> mcp::Task { + handler_entered.store(true, std::memory_order_release); + while (!server_destroyed.load(std::memory_order_acquire)) { + boost::asio::steady_timer pause(io_ctx, std::chrono::milliseconds(1)); + co_await pause.async_wait(boost::asio::use_awaitable); + } + try { + mcp::CreateMessageRequestParams request; + request.maxTokens = 1; + static_cast(co_await ctx.sample_llm(request)); + } catch (const std::exception& error) { + observed = error.what(); + } + call_returned.store(true, std::memory_order_release); + co_return nlohmann::json::object(); + }); + + boost::asio::co_spawn(io_ctx, server->run(transport, io_ctx.get_executor()), boost::asio::detached); + + auto work = boost::asio::make_work_guard(io_ctx); + std::thread runner([&io_ctx] { io_ctx.run(); }); + + ASSERT_TRUE(wait_for_condition([&] { return transport->read_calls() >= 1; }, g_teardown_budget)); + + transport->feed(make_initialize_request("init").dump()); + transport->feed(make_initialized_notification().dump()); + transport->feed(make_tool_call_request("call", "park").dump()); + + ASSERT_TRUE(wait_for_condition([&] { return handler_entered.load(std::memory_order_acquire); }, + g_teardown_budget)); + + run_with_teardown_watchdog(g_teardown_budget, [&server] { server.reset(); }); + server_destroyed.store(true, std::memory_order_release); + + EXPECT_TRUE(wait_for_condition([&] { return call_returned.load(std::memory_order_acquire); }, + g_teardown_budget)) + << "the reverse-RPC call never returned"; + EXPECT_NE(observed.find("destroyed"), std::string::npos) + << "reverse RPC after destruction reported: " << observed; + + work.reset(); + io_ctx.stop(); + runner.join(); +} diff --git a/test/server/server_elicitation_test.cpp b/test/server/server_elicitation_test.cpp index d127599..de8b886 100644 --- a/test/server/server_elicitation_test.cpp +++ b/test/server/server_elicitation_test.cpp @@ -70,8 +70,12 @@ TEST_F(ElicitationTest, FormElicitationRoundtrip) { std::vector written_messages; raw_transport->set_on_write([&written_messages, raw_transport](std::string_view msg) { - written_messages.emplace_back(msg); auto json_msg = nlohmann::json::parse(msg); + if (json_msg.value("id", "") == "initialize") { + return; + } + + written_messages.emplace_back(msg); if (json_msg.contains("method") && json_msg["method"] == "elicitation/create") { auto request_id = json_msg["id"].get(); @@ -101,6 +105,8 @@ TEST_F(ElicitationTest, FormElicitationRoundtrip) { tool_call_request["method"] = "tools/call"; tool_call_request["params"] = {{"name", "ask_user"}, {"arguments", {{"text", "hello"}}}}; + raw_transport->enqueue_message(make_initialize_request("initialize").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tool_call_request.dump()); boost::asio::co_spawn( @@ -154,8 +160,12 @@ TEST_F(ElicitationTest, URLElicitationRoundtrip) { std::vector written_messages; raw_transport->set_on_write([&written_messages, raw_transport](std::string_view msg) { - written_messages.emplace_back(msg); auto json_msg = nlohmann::json::parse(msg); + if (json_msg.value("id", "") == "initialize") { + return; + } + + written_messages.emplace_back(msg); if (json_msg.contains("method") && json_msg["method"] == "elicitation/create") { auto request_id = json_msg["id"].get(); @@ -185,6 +195,8 @@ TEST_F(ElicitationTest, URLElicitationRoundtrip) { tool_call_request["method"] = "tools/call"; tool_call_request["params"] = {{"name", "oauth_tool"}, {"arguments", {{"text", "go"}}}}; + raw_transport->enqueue_message(make_initialize_request("initialize").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tool_call_request.dump()); boost::asio::co_spawn( @@ -241,8 +253,12 @@ TEST_F(ElicitationTest, ClientDeclinesElicitation) { std::vector written_messages; raw_transport->set_on_write([&written_messages, raw_transport](std::string_view msg) { - written_messages.emplace_back(msg); auto json_msg = nlohmann::json::parse(msg); + if (json_msg.value("id", "") == "initialize") { + return; + } + + written_messages.emplace_back(msg); if (json_msg.contains("method") && json_msg["method"] == "elicitation/create") { auto request_id = json_msg["id"].get(); @@ -268,6 +284,8 @@ TEST_F(ElicitationTest, ClientDeclinesElicitation) { tool_call_request["method"] = "tools/call"; tool_call_request["params"] = {{"name", "needs_input"}, {"arguments", {{"text", "test"}}}}; + raw_transport->enqueue_message(make_initialize_request("initialize").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); raw_transport->enqueue_message(tool_call_request.dump()); boost::asio::co_spawn( diff --git a/test/server/server_handlers_test.cpp b/test/server/server_handlers_test.cpp index aa8f824..5c33a41 100644 --- a/test/server/server_handlers_test.cpp +++ b/test/server/server_handlers_test.cpp @@ -11,6 +11,7 @@ #include #include #include +#include struct AddParams { int augend = 0; @@ -59,6 +60,28 @@ class ServerHandlersTest : public ::testing::Test { }(), std::move(caps)) {} }; + + std::vector run_request(ServerSetup& setup, nlohmann::json request) { + std::vector responses; + setup.raw_transport->set_on_write([&responses, &setup](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 2) { + setup.raw_transport->close(); + } + }); + setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); + setup.raw_transport->enqueue_message(request.dump()); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await setup.server.run(setup.transport, io_ctx_.get_executor()); + }, + boost::asio::detached); + io_ctx_.run(); + return responses; + } }; TEST_F(ServerHandlersTest, SyncToolCallReturnsResult) { @@ -84,6 +107,7 @@ TEST_F(ServerHandlersTest, SyncToolCallReturnsResult) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -110,7 +134,8 @@ TEST_F(ServerHandlersTest, SyncToolCallReturnsResult) { auto& tool_response = responses[1]; EXPECT_EQ(tool_response["id"], "2"); ASSERT_TRUE(tool_response.contains("result")); - EXPECT_EQ(tool_response["result"]["sum"], 7); + EXPECT_EQ(tool_response["result"]["structuredContent"]["sum"], 7); + EXPECT_EQ(tool_response["result"]["content"][0]["type"], "text"); } TEST_F(ServerHandlersTest, AsyncToolCallReturnsResult) { @@ -135,6 +160,7 @@ TEST_F(ServerHandlersTest, AsyncToolCallReturnsResult) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -159,7 +185,58 @@ TEST_F(ServerHandlersTest, AsyncToolCallReturnsResult) { auto& tool_response = responses[1]; EXPECT_EQ(tool_response["id"], "2"); ASSERT_TRUE(tool_response.contains("result")); - EXPECT_EQ(tool_response["result"]["sum"], 30); + EXPECT_EQ(tool_response["result"]["structuredContent"]["sum"], 30); +} + +TEST_F(ServerHandlersTest, DomainObjectWithContentFieldIsStillNormalized) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool( + "domain-content", "Returns a domain object", nlohmann::json{{"type", "object"}}, + [](nlohmann::json) { + return nlohmann::json{{"content", "not an MCP content array"}, {"answer", 42}}; + }); + + auto responses = run_request( + setup, nlohmann::json{ + {"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "tools/call"}, + {"params", {{"name", "domain-content"}, {"arguments", nlohmann::json::object()}}}}); + + ASSERT_EQ(responses.size(), 2); + const auto& result = responses[1]["result"]; + EXPECT_EQ(result["structuredContent"]["content"], "not an MCP content array"); + EXPECT_EQ(result["structuredContent"]["answer"], 42); + ASSERT_TRUE(result["content"].is_array()); + EXPECT_EQ(result["content"][0]["type"], "text"); +} + +TEST_F(ServerHandlersTest, TypedHandlerExceptionBecomesProtocolToolError) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool( + "typed-failure", "Always fails", nlohmann::json{{"type", "object"}}, + [](nlohmann::json) -> nlohmann::json { throw std::runtime_error("typed failure"); }); + + auto responses = run_request( + setup, nlohmann::json{ + {"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "tools/call"}, + {"params", {{"name", "typed-failure"}, {"arguments", nlohmann::json::object()}}}}); + + ASSERT_EQ(responses.size(), 2); + const auto& result = responses[1]["result"]; + EXPECT_TRUE(result["isError"].get()); + ASSERT_TRUE(result["content"].is_array()); + EXPECT_EQ(result["content"][0]["type"], "text"); + EXPECT_NE(result["content"][0]["text"].get().find("typed failure"), std::string::npos); + EXPECT_FALSE(result.contains("structuredContent")); } TEST_F(ServerHandlersTest, ToolWithContextLogsMessage) { @@ -186,6 +263,7 @@ TEST_F(ServerHandlersTest, ToolWithContextLogsMessage) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -218,7 +296,7 @@ TEST_F(ServerHandlersTest, ToolWithContextLogsMessage) { auto& tool_response = responses[2]; EXPECT_EQ(tool_response["id"], "2"); ASSERT_TRUE(tool_response.contains("result")); - EXPECT_EQ(tool_response["result"]["sum"], 11); + EXPECT_EQ(tool_response["result"]["structuredContent"]["sum"], 11); } TEST_F(ServerHandlersTest, ToolsListReturnsRegisteredTools) { @@ -244,6 +322,7 @@ TEST_F(ServerHandlersTest, ToolsListReturnsRegisteredTools) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -271,6 +350,35 @@ TEST_F(ServerHandlersTest, ToolsListReturnsRegisteredTools) { EXPECT_EQ(tools[0]["description"], "Adds two numbers"); } +TEST_F(ServerHandlersTest, ToolsListReturnsToolsInRegistrationOrder) { + mcp::ServerCapabilities caps; + mcp::ServerCapabilities::ToolsCapability tools_cap; + caps.tools = std::move(tools_cap); + + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::Tool schema; + schema.inputSchema = nlohmann::json{{"type", "object"}}; + + for (const auto& name : {"charlie", "alpha", "bravo"}) { + mcp::Tool tool = schema; + tool.name = name; + setup.server.add_raw_tool(tool, [](const nlohmann::json&) -> nlohmann::json { + return mcp::make_tool_text_result(""); + }); + } + + auto responses = + run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, {"id", "2"}, {"method", "tools/list"}}); + + ASSERT_EQ(responses.size(), 2); + auto& tools = responses[1]["result"]["tools"]; + ASSERT_EQ(tools.size(), 3); + EXPECT_EQ(tools[0]["name"], "charlie"); + EXPECT_EQ(tools[1]["name"], "alpha"); + EXPECT_EQ(tools[2]["name"], "bravo"); +} + TEST_F(ServerHandlersTest, ResourcesReadReturnsContent) { mcp::ServerCapabilities caps; mcp::ServerCapabilities::ResourcesCapability resources_cap; @@ -304,6 +412,7 @@ TEST_F(ServerHandlersTest, ResourcesReadReturnsContent) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json read_req; read_req["jsonrpc"] = "2.0"; @@ -359,6 +468,7 @@ TEST_F(ServerHandlersTest, ResourcesListReturnsRegisteredResources) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -407,6 +517,7 @@ TEST_F(ServerHandlersTest, ResourceTemplatesListReturnsRegisteredTemplates) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -432,6 +543,351 @@ TEST_F(ServerHandlersTest, ResourceTemplatesListReturnsRegisteredTemplates) { EXPECT_EQ(templates[0]["name"], "file-access"); } +TEST_F(ServerHandlersTest, ResourceTemplateHandlerServesMatchingUri) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate tmpl; + tmpl.uriTemplate = "file:///{name}"; + tmpl.name = "files"; + setup.server.add_resource_template( + tmpl, [](mcp::ReadResourceRequestParams params) -> mcp::ReadResourceResult { + mcp::TextResourceContents content; + content.uri = params.uri; + content.text = "template content"; + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(content)); + return result; + }); + + auto responses = run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", "file:///notes.txt"}}}}); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["result"]["contents"][0]["uri"], "file:///notes.txt"); + EXPECT_EQ(responses[1]["result"]["contents"][0]["text"], "template content"); +} + +TEST_F(ServerHandlersTest, ExactResourceTakesPrecedenceOverTemplate) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate tmpl; + tmpl.uriTemplate = "file:///{name}"; + tmpl.name = "files"; + setup.server.add_resource_template( + tmpl, [](mcp::ReadResourceRequestParams params) -> mcp::ReadResourceResult { + mcp::TextResourceContents content; + content.uri = params.uri; + content.text = "template"; + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(content)); + return result; + }); + + mcp::Resource exact; + exact.uri = "file:///exact.txt"; + exact.name = "exact"; + setup.server.add_resource( + exact, [](mcp::ReadResourceRequestParams params) -> mcp::ReadResourceResult { + mcp::TextResourceContents content; + content.uri = params.uri; + content.text = "exact"; + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(content)); + return result; + }); + + auto responses = run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", "file:///exact.txt"}}}}); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["result"]["contents"][0]["text"], "exact"); +} + +TEST_F(ServerHandlersTest, ResourceTemplateRegistrationRejectsDuplicateAndEquivalentPatterns) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate first; + first.uriTemplate = "file:///{path}"; + first.name = "first"; + setup.server.add_resource_template(first); + + auto duplicate = first; + duplicate.name = "duplicate"; + EXPECT_THROW(setup.server.add_resource_template(duplicate), std::invalid_argument); + + mcp::ResourceTemplate equivalent; + equivalent.uriTemplate = "file:///{name}"; + equivalent.name = "equivalent"; + EXPECT_THROW(setup.server.add_resource_template(equivalent), std::invalid_argument); + + mcp::ResourceTemplate malformed; + malformed.uriTemplate = "file:///{path"; + malformed.name = "malformed"; + EXPECT_THROW(setup.server.add_resource_template(malformed), std::invalid_argument); +} + +TEST_F(ServerHandlersTest, AmbiguousResourceTemplateMatchReturnsInvalidParams) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + auto register_template = [&setup](std::string uri_template, std::string name) { + mcp::ResourceTemplate tmpl; + tmpl.uriTemplate = std::move(uri_template); + tmpl.name = std::move(name); + setup.server.add_resource_template( + tmpl, [](mcp::ReadResourceRequestParams) { return mcp::ReadResourceResult{}; }); + }; + register_template("file:///{+path}", "all-files"); + register_template("file:///fixed/{name}", "fixed-files"); + + auto responses = + run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", "file:///fixed/item.txt"}}}}); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_INVALID_PARAMS); + EXPECT_NE(responses[1]["error"]["message"].get().find("Ambiguous"), std::string::npos); +} + +TEST_F(ServerHandlersTest, MetadataOnlyTemplateDoesNotBlockHandledTemplate) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate metadata_only; + metadata_only.uriTemplate = "file:///{+path}"; + metadata_only.name = "all-files"; + setup.server.add_resource_template(metadata_only); + + mcp::ResourceTemplate handled; + handled.uriTemplate = "file:///fixed/{name}"; + handled.name = "fixed-files"; + setup.server.add_resource_template( + handled, [](mcp::ReadResourceRequestParams params) { + mcp::TextResourceContents content; + content.uri = params.uri; + content.text = "handled"; + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(content)); + return result; + }); + + auto responses = + run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", "file:///fixed/item.txt"}}}}); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["result"]["contents"][0]["text"], "handled"); +} + +TEST_F(ServerHandlersTest, UnknownResourceReturnsInvalidParams) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + auto responses = + run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", "file:///does-not-exist.txt"}}}}); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_INVALID_PARAMS); +} + +TEST_F(ServerHandlersTest, UnknownResourceTemplateReturnsInvalidParams) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate tmpl; + tmpl.uriTemplate = "file:///docs/{name}"; + tmpl.name = "docs"; + setup.server.add_resource_template( + tmpl, [](mcp::ReadResourceRequestParams) { return mcp::ReadResourceResult{}; }); + + auto responses = + run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", "file:///other/item.txt"}}}}); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_INVALID_PARAMS); +} + +TEST_F(ServerHandlersTest, OverlongUriIsRejectedWithoutReachingTemplateMatching) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate tmpl; + tmpl.uriTemplate = "file:///{+path}"; + tmpl.name = "files"; + setup.server.add_resource_template( + tmpl, [](mcp::ReadResourceRequestParams) { return mcp::ReadResourceResult{}; }); + + const std::string uri = "file:///" + std::string(64 * 1024, 'a'); + auto responses = run_request( + setup, + nlohmann::json{ + {"jsonrpc", "2.0"}, {"id", "2"}, {"method", "resources/read"}, {"params", {{"uri", uri}}}}); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("error")); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_INVALID_PARAMS); + // Says the URI was too long, rather than that no such resource exists, and does not echo it. + const auto message = responses[1]["error"]["message"].get(); + EXPECT_NE(message.find("limit"), std::string::npos); + EXPECT_EQ(message.find("Unknown resource"), std::string::npos); + EXPECT_LT(message.size(), uri.size()); +} + +// A URI at the accepted limit is served, so the bound admits what it claims to admit. The match is +// run on a std::thread, which carries the platform's default stack — the configuration the bound is +// sized for. On gtest's main thread (8 MB on Linux and macOS) this would prove nothing. +TEST_F(ServerHandlersTest, UriAtTheLengthLimitIsStillMatchedByTemplate) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate tmpl; + tmpl.uriTemplate = "file:///{+path}"; + tmpl.name = "files"; + setup.server.add_resource_template( + tmpl, [](mcp::ReadResourceRequestParams params) { + mcp::TextResourceContents content; + content.uri = params.uri; + content.text = "matched"; + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(content)); + return result; + }); + + constexpr std::size_t limit = 512; + const std::string prefix = "file:///"; + const std::string uri = prefix + std::string(limit - prefix.size(), 'a'); + ASSERT_EQ(uri.size(), limit); + + std::vector responses; + setup.raw_transport->set_on_write([&responses, &setup](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 2) { + setup.raw_transport->close(); + } + }); + setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); + setup.raw_transport->enqueue_message(nlohmann::json{ + {"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", uri}}}}.dump()); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await setup.server.run(setup.transport, io_ctx_.get_executor()); + }, + boost::asio::detached); + + std::thread pump([this]() { io_ctx_.run(); }); + pump.join(); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_EQ(responses[1]["result"]["contents"][0]["text"], "matched"); +} + +TEST_F(ServerHandlersTest, ExactResourceIsServedRegardlessOfUriLength) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::Resource exact; + exact.uri = "file:///" + std::string(64 * 1024, 'a'); + exact.name = "long"; + setup.server.add_resource( + exact, [](mcp::ReadResourceRequestParams params) { + mcp::TextResourceContents content; + content.uri = params.uri; + content.text = "exact"; + mcp::ReadResourceResult result; + result.contents.emplace_back(std::move(content)); + return result; + }); + + auto responses = run_request(setup, nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "resources/read"}, + {"params", {{"uri", exact.uri}}}}); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["result"]["contents"][0]["text"], "exact"); +} + +TEST_F(ServerHandlersTest, AdjacentTemplateExpressionsAreRejectedAtRegistration) { + mcp::ServerCapabilities caps; + caps.resources = mcp::ServerCapabilities::ResourcesCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::ResourceTemplate adjacent; + adjacent.uriTemplate = "file:///{dir}{name}"; + adjacent.name = "adjacent"; + EXPECT_THROW(setup.server.add_resource_template(adjacent), std::invalid_argument); + + mcp::ResourceTemplate separated; + separated.uriTemplate = "file:///{dir}/{name}"; + separated.name = "separated"; + EXPECT_NO_THROW(setup.server.add_resource_template(separated)); +} + +TEST_F(ServerHandlersTest, DuplicateNamedRegistrationsAreRejected) { + ServerSetup setup(io_ctx_, mcp::ServerCapabilities{}); + + auto tool_handler = [](nlohmann::json) { return nlohmann::json::object(); }; + setup.server.add_tool( + "duplicate", "first", nlohmann::json{{"type", "object"}}, tool_handler); + EXPECT_THROW((setup.server.add_tool( + "duplicate", "second", nlohmann::json{{"type", "object"}}, tool_handler)), + std::invalid_argument); + + mcp::Resource resource; + resource.uri = "file:///duplicate"; + resource.name = "first"; + auto resource_handler = [](mcp::ReadResourceRequestParams) { return mcp::ReadResourceResult{}; }; + setup.server.add_resource( + resource, resource_handler); + resource.name = "second"; + EXPECT_THROW((setup.server.add_resource( + resource, resource_handler)), + std::invalid_argument); + + mcp::Prompt prompt; + prompt.name = "duplicate"; + auto prompt_handler = [](mcp::GetPromptRequestParams) { return mcp::GetPromptResult{}; }; + setup.server.add_prompt(prompt, prompt_handler); + prompt.description = "second"; + EXPECT_THROW((setup.server.add_prompt( + prompt, prompt_handler)), + std::invalid_argument); +} + TEST_F(ServerHandlersTest, PromptsGetReturnsPromptMessages) { mcp::ServerCapabilities caps; mcp::ServerCapabilities::PromptsCapability prompts_cap; @@ -479,6 +935,7 @@ TEST_F(ServerHandlersTest, PromptsGetReturnsPromptMessages) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json get_req; get_req["jsonrpc"] = "2.0"; @@ -532,6 +989,7 @@ TEST_F(ServerHandlersTest, PromptsListReturnsRegisteredPrompts) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -573,6 +1031,7 @@ TEST_F(ServerHandlersTest, UnknownToolReturnsError) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -629,6 +1088,7 @@ TEST_F(ServerHandlersTest, ToolsListPaginationFirstPage) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -679,6 +1139,7 @@ TEST_F(ServerHandlersTest, ToolsListPaginationWithCursor) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -729,6 +1190,7 @@ TEST_F(ServerHandlersTest, ToolsListPaginationLastPage) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -776,6 +1238,7 @@ TEST_F(ServerHandlersTest, PaginationDisabledByDefault) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -828,6 +1291,7 @@ TEST_F(ServerHandlersTest, ResourcesListPagination) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -880,6 +1344,7 @@ TEST_F(ServerHandlersTest, ToolWithOutputSchemaIncludesStructuredContent) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -900,12 +1365,39 @@ TEST_F(ServerHandlersTest, ToolWithOutputSchemaIncludesStructuredContent) { ASSERT_EQ(responses.size(), 2); auto& result = responses[1]["result"]; - EXPECT_EQ(result["sum"], 30); ASSERT_TRUE(result.contains("structuredContent")); EXPECT_EQ(result["structuredContent"]["sum"], 30); + EXPECT_EQ(result["content"][0]["text"], R"({"sum":30})"); } -TEST_F(ServerHandlersTest, ToolWithoutOutputSchemaNoStructuredContent) { +TEST_F(ServerHandlersTest, ToolWithOutputSchemaCanReturnAnErrorWithoutStructuredContent) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool( + "failing_structured", "Fails before producing structured output", + nlohmann::json{{"type", "object"}}, nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> nlohmann::json { + throw std::runtime_error("structured tool failed"); + }); + + auto responses = run_request( + setup, + nlohmann::json{ + {"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "tools/call"}, + {"params", {{"name", "failing_structured"}, {"arguments", nlohmann::json::object()}}}}); + + ASSERT_EQ(responses.size(), 2); + const auto& result = responses[1]["result"]; + EXPECT_TRUE(result["isError"].get()); + EXPECT_FALSE(result.contains("structuredContent")); + EXPECT_EQ(result["content"][0]["text"], "structured tool failed"); +} + +TEST_F(ServerHandlersTest, ToolWithoutOutputSchemaStillReturnsStructuredContent) { mcp::ServerCapabilities caps; mcp::ServerCapabilities::ToolsCapability tools_cap; caps.tools = std::move(tools_cap); @@ -925,6 +1417,7 @@ TEST_F(ServerHandlersTest, ToolWithoutOutputSchemaNoStructuredContent) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -945,8 +1438,8 @@ TEST_F(ServerHandlersTest, ToolWithoutOutputSchemaNoStructuredContent) { ASSERT_EQ(responses.size(), 2); auto& result = responses[1]["result"]; - EXPECT_EQ(result["sum"], 8); - EXPECT_FALSE(result.contains("structuredContent")); + EXPECT_EQ(result["structuredContent"]["sum"], 8); + EXPECT_EQ(result["content"][0]["type"], "text"); } TEST_F(ServerHandlersTest, ToolsListIncludesOutputSchema) { @@ -972,6 +1465,7 @@ TEST_F(ServerHandlersTest, ToolsListIncludesOutputSchema) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -1022,6 +1516,7 @@ TEST_F(ServerHandlersTest, CompletionReturnsResults) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json complete_req; complete_req["jsonrpc"] = "2.0"; @@ -1066,13 +1561,14 @@ TEST_F(ServerHandlersTest, CompletionWithoutHandlerReturnsError) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json complete_req; complete_req["jsonrpc"] = "2.0"; complete_req["id"] = "2"; complete_req["method"] = "completion/complete"; complete_req["params"] = nlohmann::json{{"ref", {{"type", "ref/prompt"}, {"name", "test"}}}, - {"argument", {{"name", "arg", "value", "val"}}}}; + {"argument", {{"name", "arg"}, {"value", "val"}}}}; setup.raw_transport->enqueue_message(complete_req.dump()); boost::asio::co_spawn( @@ -1089,3 +1585,274 @@ TEST_F(ServerHandlersTest, CompletionWithoutHandlerReturnsError) { ASSERT_TRUE(error_response.contains("error")); EXPECT_EQ(error_response["error"]["code"], mcp::g_METHOD_NOT_FOUND); } + +// A tool name is chosen entirely by the client, and an unknown one is the ordinary error path of +// any `tools/call`. The name reaches a diagnostic that the server operator reads, so a name +// carrying CR/LF forges a line in whatever that diagnostic lands in, and a bidi override reorders +// the rest of it. Same defect as the peer-controlled text in the auth diagnostics, opposite +// direction: here the untrusted side is the client. +TEST_F(ServerHandlersTest, UnknownToolDiagnosticFlattensClientChosenName) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + // "zqtripwire" is a token no other code path produces; see the tripwire assertions below. + // \xe2\x80\xae is U+202E RIGHT-TO-LEFT OVERRIDE, written escaped so this source file does not + // itself contain a bidi override. + const std::string forged_name = + "zqtripwire\r\n2026-09-20 ERROR forged line from the client\xe2\x80\xae reordered tail"; + + nlohmann::json call_req; + call_req["jsonrpc"] = "2.0"; + call_req["id"] = "2"; + call_req["method"] = "tools/call"; + call_req["params"] = nlohmann::json{{"name", forged_name}, {"arguments", nlohmann::json::object()}}; + + auto responses = run_request(setup, std::move(call_req)); + + ASSERT_EQ(responses.size(), 2); + auto& error_response = responses[1]; + ASSERT_TRUE(error_response.contains("error")) << "response: " << error_response.dump(); + const auto message = error_response["error"]["message"].get(); + + // Tripwire. The payload must actually have reached the "Unknown tool" site rather than being + // rejected earlier by params validation or routed to some other diagnostic that flattens its + // own message; either would make the assertions below pass for a reason unrelated to the site + // under test. The code pins WHICH of the two "Unknown tool" sites this is: the RPC path refuses + // in handle_tools_call_wire with g_METHOD_NOT_FOUND, whereas the invoke_tool path throws and + // surfaces as g_INTERNAL_ERROR through the dispatcher. Dump the message on failure so a vacuous + // pass cannot hide. + ASSERT_EQ(error_response["error"]["code"], mcp::g_METHOD_NOT_FOUND) + << "response: " << error_response.dump(); + ASSERT_EQ(message.rfind("Unknown tool: ", 0), 0U) << "actual message: " << message; + ASSERT_NE(message.find("zqtripwire"), std::string::npos) << "actual message: " << message; + + EXPECT_EQ(message.find('\r'), std::string::npos) << "actual message: " << message; + EXPECT_EQ(message.find('\n'), std::string::npos) << "actual message: " << message; + EXPECT_EQ(message.find("\xe2\x80\xae"), std::string::npos) << "actual message: " << message; +} + +// The same site must also bound the name, so a client cannot flood the operator's log through a +// megabyte-long tool name. +TEST_F(ServerHandlersTest, UnknownToolDiagnosticBoundsClientChosenName) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + const std::string forged_name = "zqtripwire" + std::string(64 * 1024, 'A'); + + nlohmann::json call_req; + call_req["jsonrpc"] = "2.0"; + call_req["id"] = "2"; + call_req["method"] = "tools/call"; + call_req["params"] = nlohmann::json{{"name", forged_name}, {"arguments", nlohmann::json::object()}}; + + auto responses = run_request(setup, std::move(call_req)); + + ASSERT_EQ(responses.size(), 2); + auto& error_response = responses[1]; + ASSERT_TRUE(error_response.contains("error")) << "response: " << error_response.dump(); + const auto message = error_response["error"]["message"].get(); + + ASSERT_EQ(error_response["error"]["code"], mcp::g_METHOD_NOT_FOUND) + << "actual message prefix: " << message.substr(0, 64); + ASSERT_EQ(message.rfind("Unknown tool: ", 0), 0U) + << "actual message prefix: " << message.substr(0, 64); + ASSERT_NE(message.find("zqtripwire"), std::string::npos) + << "actual message prefix: " << message.substr(0, 64); + EXPECT_LT(message.size(), forged_name.size()) << "message size: " << message.size(); + EXPECT_LE(message.size(), std::size_t{512}) << "message size: " << message.size(); +} + +// The second "Unknown tool" site. `invoke_tool` bypasses the JSON-RPC loop for json_only +// deployments, so the name it is given is whatever the embedding passes in -- peer-chosen text in +// exactly the deployments that use this entry point. It is a distinct site from the one the RPC +// path takes, and it throws rather than building an error frame, so it needs its own test. +TEST_F(ServerHandlersTest, InvokeToolUnknownNameDiagnosticIsFlattened) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + const std::string forged_name = + "zqtripwire\r\n2026-09-20 ERROR forged line from the caller\xe2\x80\xae reordered tail"; + + std::string message; + bool threw = false; + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + try { + static_cast( + co_await setup.server.invoke_tool(forged_name, nlohmann::json::object())); + } catch (const std::exception& error) { + threw = true; + message = error.what(); + } + }, + boost::asio::detached); + io_ctx_.run(); + + // Tripwire: the name must have reached the invoke_tool site itself, not some earlier refusal. + ASSERT_TRUE(threw) << "invoke_tool accepted an unknown tool name"; + ASSERT_EQ(message.rfind("Unknown tool: ", 0), 0U) << "actual message: " << message; + ASSERT_NE(message.find("zqtripwire"), std::string::npos) << "actual message: " << message; + + EXPECT_EQ(message.find('\r'), std::string::npos) << "actual message: " << message; + EXPECT_EQ(message.find('\n'), std::string::npos) << "actual message: " << message; + EXPECT_EQ(message.find("\xe2\x80\xae"), std::string::npos) << "actual message: " << message; +} + +// --- Explicit-null request members --- +// +// An absent JSON member and one serialized as explicit `null` are the same statement on the +// wire. Both requests below are messages a conforming peer may send; neither may be rejected. + +TEST_F(ServerHandlersTest, PromptsGetAcceptsExplicitNullArguments) { + mcp::ServerCapabilities caps; + mcp::ServerCapabilities::PromptsCapability prompts_cap; + caps.prompts = std::move(prompts_cap); + + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::Prompt prompt; + prompt.name = "greeting"; + setup.server.add_prompt( + std::move(prompt), [](mcp::GetPromptRequestParams params) -> mcp::GetPromptResult { + EXPECT_FALSE(params.arguments.has_value()); + + mcp::TextContent text; + text.text = "Hello, World!"; + + mcp::PromptMessage msg; + msg.role = mcp::Role::eUser; + msg.content = std::move(text); + + mcp::GetPromptResult result; + result.messages.push_back(std::move(msg)); + return result; + }); + + std::vector responses; + setup.raw_transport->set_on_write([&responses, &setup](std::string_view msg) { + responses.push_back(nlohmann::json::parse(msg)); + if (responses.size() == 2) { + setup.raw_transport->close(); + } + }); + + setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); + + nlohmann::json get_req; + get_req["jsonrpc"] = "2.0"; + get_req["id"] = "2"; + get_req["method"] = "prompts/get"; + get_req["params"] = nlohmann::json{{"name", "greeting"}, {"arguments", nullptr}}; + setup.raw_transport->enqueue_message(get_req.dump()); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await setup.server.run(setup.transport, io_ctx_.get_executor()); + }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 2); + auto& get_response = responses[1]; + EXPECT_EQ(get_response["id"], "2"); + ASSERT_FALSE(get_response.contains("error")) << "actual error: " << get_response["error"].dump(); + ASSERT_TRUE(get_response.contains("result")); + auto& messages = get_response["result"]["messages"]; + ASSERT_EQ(messages.size(), 1); + EXPECT_EQ(messages[0]["content"]["text"], "Hello, World!"); +} + +TEST_F(ServerHandlersTest, ToolsListPaginationAcceptsExplicitNullCursor) { + // Only reachable with a page size set and a non-empty collection: paginate() returns the + // whole slice before it ever looks at the cursor when either is absent. + mcp::ServerCapabilities caps; + mcp::ServerCapabilities::ToolsCapability tools_cap; + caps.tools = std::move(tools_cap); + + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.set_page_size(2); + + for (int i = 0; i < 5; ++i) { + auto name = "tool_" + std::to_string(i); + setup.server.add_tool( + name, "Tool " + std::to_string(i), nlohmann::json{{"type", "object"}}, + [](AddParams p) -> AddResult { return AddResult{p.augend + p.addend}; }); + } + + std::vector responses; + setup.raw_transport->set_on_write([&responses, &setup](std::string_view msg) { + responses.push_back(nlohmann::json::parse(msg)); + if (responses.size() == 2) { + setup.raw_transport->close(); + } + }); + + setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); + + nlohmann::json list_req; + list_req["jsonrpc"] = "2.0"; + list_req["id"] = "2"; + list_req["method"] = "tools/list"; + list_req["params"] = nlohmann::json{{"cursor", nullptr}}; + setup.raw_transport->enqueue_message(list_req.dump()); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await setup.server.run(setup.transport, io_ctx_.get_executor()); + }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 2); + auto& list_response = responses[1]; + ASSERT_FALSE(list_response.contains("error")) << "actual error: " << list_response["error"].dump(); + auto& result = list_response["result"]; + ASSERT_EQ(result["tools"].size(), 2); + EXPECT_EQ(result["tools"][0]["name"], "tool_0"); + ASSERT_TRUE(result.contains("nextCursor")); + EXPECT_EQ(result["nextCursor"], "2"); +} + +TEST_F(ServerHandlersTest, InitializeAcceptsExplicitNullClientInfoOptionals) { + // clientInfo is an Implementation, so this breaks the handshake on a peer's first message. + ServerSetup setup(io_ctx_, mcp::ServerCapabilities{}); + + std::vector responses; + setup.raw_transport->set_on_write([&responses, &setup](std::string_view msg) { + responses.push_back(nlohmann::json::parse(msg)); + setup.raw_transport->close(); + }); + + auto init_req = make_initialize_request("1"); + init_req["params"]["clientInfo"]["title"] = nullptr; + init_req["params"]["clientInfo"]["description"] = nullptr; + init_req["params"]["clientInfo"]["websiteUrl"] = nullptr; + init_req["params"]["clientInfo"]["icons"] = nullptr; + setup.raw_transport->enqueue_message(init_req.dump()); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await setup.server.run(setup.transport, io_ctx_.get_executor()); + }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 1); + auto& init_response = responses[0]; + ASSERT_FALSE(init_response.contains("error")) << "actual error: " << init_response["error"].dump(); + ASSERT_TRUE(init_response.contains("result")); + EXPECT_EQ(init_response["result"]["protocolVersion"], std::string(mcp::g_LATEST_PROTOCOL_VERSION)); +} diff --git a/test/server/server_http_test.cpp b/test/server/server_http_test.cpp index 18e5c1c..0c48860 100644 --- a/test/server/server_http_test.cpp +++ b/test/server/server_http_test.cpp @@ -43,6 +43,10 @@ nlohmann::json make_tool_call_request(std::string_view id, const std::string& to {"params", {{"name", tool_name}, {"arguments", args}}}}; } +nlohmann::json make_initialized_notification() { + return {{"jsonrpc", "2.0"}, {"method", "notifications/initialized"}}; +} + nlohmann::json greet_schema() { return { {"type", "object"}, @@ -119,6 +123,7 @@ TEST_F(ServerHttpTest, RunHttpInitializesAndResponds) { session_id = std::string(session_it->value()); } + static_cast(send_json_rpc(stream, make_initialized_notification(), session_id)); auto tool_resp = send_json_rpc(stream, make_tool_call_request("2", "greet", {{"name", "HTTP"}}), session_id); auto tool_json = nlohmann::json::parse(tool_resp.body()); @@ -135,7 +140,7 @@ TEST_F(ServerHttpTest, RunHttpInitializesAndResponds) { EXPECT_EQ(tool_json["id"], "2"); ASSERT_TRUE(tool_json.contains("result")); - EXPECT_EQ(tool_json["result"]["greeting"], "Hello, HTTP"); + EXPECT_EQ(tool_json["result"]["structuredContent"]["greeting"], "Hello, HTTP"); } TEST_F(ServerHttpTest, RunHttpInvalidAddressThrows) { diff --git a/test/server/server_logging_test.cpp b/test/server/server_logging_test.cpp index 889aff8..6654f5f 100644 --- a/test/server/server_logging_test.cpp +++ b/test/server/server_logging_test.cpp @@ -60,6 +60,7 @@ TEST_F(ServerLoggingTest, SetLevelChangesServerLogLevel) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json set_level_req; set_level_req["jsonrpc"] = "2.0"; @@ -115,6 +116,7 @@ TEST_F(ServerLoggingTest, LogFilteredByLevel) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); // Set level to Warning (filters out Info and Debug) nlohmann::json set_level_req; @@ -185,6 +187,7 @@ TEST_F(ServerLoggingTest, LogWithLoggerField) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -240,6 +243,7 @@ TEST_F(ServerLoggingTest, ListChangedNotificationsSent) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; diff --git a/test/server/server_middleware_test.cpp b/test/server/server_middleware_test.cpp index 34a6473..caef490 100644 --- a/test/server/server_middleware_test.cpp +++ b/test/server/server_middleware_test.cpp @@ -12,6 +12,7 @@ #include #include #include +#include #include #include @@ -63,6 +64,7 @@ TEST(ServerMiddlewareTest, MiddlewareModifiesParams) { }); transport_ptr->enqueue_message(make_initialize_request("1").dump()); + transport_ptr->enqueue_message(make_initialized_notification().dump()); json call_req{{"jsonrpc", "2.0"}, {"id", "2"}, @@ -79,7 +81,7 @@ TEST(ServerMiddlewareTest, MiddlewareModifiesParams) { return msg.contains("id") && msg["id"] == "2" && msg.contains("result"); }); ASSERT_NE(it, responses.end()); - EXPECT_EQ((*it)["result"]["sum"], 115); + EXPECT_EQ((*it)["result"]["structuredContent"]["sum"], 115); } // Test 2: Middleware short-circuits (returns without calling next) @@ -99,7 +101,7 @@ TEST(ServerMiddlewareTest, MiddlewareShortCircuits) { // Middleware that blocks "forbidden" tool without calling next server.use([](Context& ctx, const json& params, TypeErasedHandler next) -> Task { if (params.contains("name") && params["name"] == "forbidden") { - co_return json{{"error", "Tool 'forbidden' is not allowed"}}; + throw std::runtime_error("Tool 'forbidden' is not allowed"); } co_return co_await next(ctx, params); }); @@ -124,11 +126,12 @@ TEST(ServerMiddlewareTest, MiddlewareShortCircuits) { }); transport_ptr->enqueue_message(make_initialize_request("1").dump()); + transport_ptr->enqueue_message(make_initialized_notification().dump()); json call_req{{"jsonrpc", "2.0"}, {"id", "2"}, {"method", "tools/call"}, - {"params", {{"name", "forbidden"}, {"arguments", {}}}}}; + {"params", {{"name", "forbidden"}, {"arguments", nlohmann::json::object()}}}}; transport_ptr->enqueue_message(call_req.dump()); boost::asio::co_spawn(io, server.run(transport, io.get_executor()), boost::asio::detached); @@ -137,13 +140,14 @@ TEST(ServerMiddlewareTest, MiddlewareShortCircuits) { // Verify the tool was never executed EXPECT_FALSE(tool_executed); - // Check that a result was sent with an error field + // A middleware that throws is signalling that the call must not proceed, which belongs on the + // protocol channel rather than in a successful tool result. auto it = std::find_if(responses.begin(), responses.end(), [](const json& msg) { - return msg.contains("id") && msg["id"] == "2" && msg.contains("result"); + return msg.contains("id") && msg["id"] == "2" && msg.contains("error"); }); ASSERT_NE(it, responses.end()); - EXPECT_TRUE((*it)["result"].contains("error")); - EXPECT_TRUE((*it)["result"]["error"].get().find("forbidden") != std::string::npos); + EXPECT_EQ((*it)["error"]["code"], mcp::g_INTERNAL_ERROR); + EXPECT_TRUE((*it)["error"]["message"].get().find("forbidden") != std::string::npos); } // Test 3: Multiple middlewares execute in order (1→2→handler→2→1) @@ -198,11 +202,12 @@ TEST(ServerMiddlewareTest, MultipleMiddlewaresExecuteInOrder) { }); transport_ptr->enqueue_message(make_initialize_request("1").dump()); + transport_ptr->enqueue_message(make_initialized_notification().dump()); json call_req{{"jsonrpc", "2.0"}, {"id", "2"}, {"method", "tools/call"}, - {"params", {{"name", "track"}, {"arguments", {}}}}}; + {"params", {{"name", "track"}, {"arguments", nlohmann::json::object()}}}}; transport_ptr->enqueue_message(call_req.dump()); boost::asio::co_spawn(io, server.run(transport, io.get_executor()), boost::asio::detached); @@ -265,6 +270,7 @@ TEST(ServerMiddlewareTest, MiddlewarePostProcessesResult) { }); transport_ptr->enqueue_message(make_initialize_request("1").dump()); + transport_ptr->enqueue_message(make_initialized_notification().dump()); json call_req{{"jsonrpc", "2.0"}, {"id", "2"}, @@ -280,6 +286,6 @@ TEST(ServerMiddlewareTest, MiddlewarePostProcessesResult) { return msg.contains("id") && msg["id"] == "2" && msg.contains("result"); }); ASSERT_NE(it, responses.end()); - EXPECT_EQ((*it)["result"]["sum"], 40); - EXPECT_EQ((*it)["result"]["metadata"], "processed by middleware"); + EXPECT_EQ((*it)["result"]["structuredContent"]["sum"], 40); + EXPECT_EQ((*it)["result"]["structuredContent"]["metadata"], "processed by middleware"); } diff --git a/test/server/server_progress_test.cpp b/test/server/server_progress_test.cpp index 6812f89..0dc6156 100644 --- a/test/server/server_progress_test.cpp +++ b/test/server/server_progress_test.cpp @@ -66,6 +66,7 @@ TEST_F(ServerProgressTest, ProgressNotificationSentWithToken) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -130,6 +131,7 @@ TEST_F(ServerProgressTest, ProgressSilentWithoutToken) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); // No _meta.progressToken in this request nlohmann::json call_req; @@ -185,6 +187,7 @@ TEST_F(ServerProgressTest, MultipleProgressNotifications) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; diff --git a/test/server/server_stdio_test.cpp b/test/server/server_stdio_test.cpp index 544ebe7..533a90a 100644 --- a/test/server/server_stdio_test.cpp +++ b/test/server/server_stdio_test.cpp @@ -3,10 +3,15 @@ #include +#include +#include #include #include +#include #include +#include #include +#include #include #include @@ -30,6 +35,10 @@ nlohmann::json make_tool_call_request(std::string_view id, const std::string& to {"params", {{"name", tool_name}, {"arguments", args}}}}; } +nlohmann::json make_initialized_notification() { + return {{"jsonrpc", "2.0"}, {"method", "notifications/initialized"}}; +} + nlohmann::json greet_schema() { return { {"type", "object"}, @@ -93,6 +102,62 @@ class BlockingStreambuf : public std::streambuf { bool closed_ = false; }; +/// A streambuf that collects what the server writes under a mutex and lets the test thread block +/// until whole response lines have arrived. +/// +/// Polling a std::ostringstream that the server thread is writing into would be a data race on the +/// stream's own buffer, and run_stdio() offers its caller no way to learn that a response has been +/// produced. +class CollectingStreambuf : public std::streambuf { + public: + /// Blocks until `count` newline-terminated lines have been written or the budget expires. + /// Returns whether the count was reached. + bool wait_for_lines(std::size_t count, std::chrono::milliseconds budget) { + std::unique_lock lock(mu_); + return cv_.wait_for(lock, budget, [&] { return lines_ >= count; }); + } + + std::string str() { + std::lock_guard lock(mu_); + return buffer_; + } + + protected: + int overflow(int ch) override { + if (ch == traits_type::eof()) { + return traits_type::not_eof(ch); + } + const auto byte = traits_type::to_char_type(ch); + { + std::lock_guard lock(mu_); + buffer_.push_back(byte); + if (byte == '\n') { + ++lines_; + } + } + cv_.notify_all(); + return ch; + } + + std::streamsize xsputn(const char* data, std::streamsize count) override { + { + std::lock_guard lock(mu_); + buffer_.append(data, static_cast(count)); + lines_ += static_cast(std::count(data, data + count, '\n')); + } + cv_.notify_all(); + return count; + } + + private: + std::mutex mu_; + std::condition_variable cv_; + std::string buffer_; + std::size_t lines_ = 0; +}; + +constexpr std::chrono::milliseconds g_response_budget{10000}; + } // namespace class ServerStdioTest : public ::testing::Test { @@ -120,24 +185,19 @@ class ServerStdioTest : public ::testing::Test { TEST_F(ServerStdioTest, RunStdioInitializesAndResponds) { BlockingStreambuf sbuf; std::istream input(&sbuf); - std::ostringstream output; + CollectingStreambuf obuf; + std::ostream output(&obuf); sbuf.feed(make_initialize_request("1").dump() + "\n"); std::thread server_thread([&] { server_->run_stdio(input, output); }); - // Wait for the response to appear - for (int i = 0; i < 50; ++i) { - if (!output.str().empty()) { - break; - } - std::this_thread::sleep_for(std::chrono::milliseconds(20)); - } + EXPECT_TRUE(obuf.wait_for_lines(1, g_response_budget)) << "no response line was written"; sbuf.close(); server_thread.join(); - auto responses = parse_responses(output.str()); + auto responses = parse_responses(obuf.str()); ASSERT_EQ(responses.size(), 1); EXPECT_EQ(responses[0]["id"], "1"); EXPECT_TRUE(responses[0].contains("result")); @@ -147,32 +207,27 @@ TEST_F(ServerStdioTest, RunStdioInitializesAndResponds) { TEST_F(ServerStdioTest, RunStdioHandlesToolCall) { BlockingStreambuf sbuf; std::istream input(&sbuf); - std::ostringstream output; + CollectingStreambuf obuf; + std::ostream output(&obuf); sbuf.feed(make_initialize_request("1").dump() + "\n"); + sbuf.feed(make_initialized_notification().dump() + "\n"); sbuf.feed(make_tool_call_request("2", "greet", {{"name", "World"}}).dump() + "\n"); std::thread server_thread([&] { server_->run_stdio(input, output); }); - // Wait for both responses - for (int i = 0; i < 100; ++i) { - auto lines = parse_responses(output.str()); - if (lines.size() >= 2) { - break; - } - std::this_thread::sleep_for(std::chrono::milliseconds(20)); - } + EXPECT_TRUE(obuf.wait_for_lines(2, g_response_budget)) << "fewer than two response lines"; sbuf.close(); server_thread.join(); - auto responses = parse_responses(output.str()); + auto responses = parse_responses(obuf.str()); ASSERT_GE(responses.size(), 2); auto tool_response = responses[1]; EXPECT_EQ(tool_response["id"], "2"); ASSERT_TRUE(tool_response.contains("result")); - EXPECT_EQ(tool_response["result"]["greeting"], "Hello, World"); + EXPECT_EQ(tool_response["result"]["structuredContent"]["greeting"], "Hello, World"); } TEST_F(ServerStdioTest, RunStdioExitsCleanlyOnEmptyInput) { @@ -184,6 +239,9 @@ TEST_F(ServerStdioTest, RunStdioExitsCleanlyOnEmptyInput) { EXPECT_TRUE(output.str().empty()); } +// Returning is the whole assertion: what this detects is run_stdio() failing to shut down on a +// signal, which shows up as a join that never completes rather than as a wrong value. The absence +// of an EXPECT here is deliberate. TEST_F(ServerStdioTest, RunStdioSignalCausesShutdown) { BlockingStreambuf sbuf; std::istream input(&sbuf); diff --git a/test/server/server_subscriptions_test.cpp b/test/server/server_subscriptions_test.cpp index a52ca23..e9d7abf 100644 --- a/test/server/server_subscriptions_test.cpp +++ b/test/server/server_subscriptions_test.cpp @@ -41,6 +41,7 @@ TEST_F(ServerSubscriptionsTest, SubscribeReturnsEmptyResult) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json sub_req; sub_req["jsonrpc"] = "2.0"; @@ -86,6 +87,7 @@ TEST_F(ServerSubscriptionsTest, UnsubscribeReturnsEmptyResult) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json sub_req; sub_req["jsonrpc"] = "2.0"; @@ -157,6 +159,7 @@ TEST_F(ServerSubscriptionsTest, NotifyResourceUpdatedSendsToSubscribers) { }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); // Subscribe to a resource nlohmann::json sub_req; @@ -230,6 +233,7 @@ TEST_F(ServerSubscriptionsTest, NotifyResourceUpdatedNotSentForUnsubscribedUri) }); raw_transport->enqueue_message(make_initialize_request("1").dump()); + raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; diff --git a/test/server/server_tool_test.cpp b/test/server/server_tool_test.cpp index 8c1d0a8..11fd83c 100644 --- a/test/server/server_tool_test.cpp +++ b/test/server/server_tool_test.cpp @@ -7,8 +7,11 @@ #include #include #include +#include +#include #include #include +#include #include #include @@ -33,6 +36,33 @@ class ServerToolTest : public ::testing::Test { }(), std::move(caps)) {} }; + + std::vector call_tool(ServerSetup& setup, std::string name) { + std::vector responses; + setup.raw_transport->set_on_write([&responses, &setup](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 2) { + setup.raw_transport->close(); + } + }); + setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); + setup.raw_transport->enqueue_message(nlohmann::json{ + {"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "tools/call"}, + {"params", {{"name", std::move(name)}, {"arguments", nlohmann::json::object()}}}} + .dump()); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await setup.server.run(setup.transport, io_ctx_.get_executor()); + }, + boost::asio::detached); + io_ctx_.run(); + return responses; + } }; TEST_F(ServerToolTest, NonTemplateSyncHandlerReturnsResult) { @@ -56,6 +86,7 @@ TEST_F(ServerToolTest, NonTemplateSyncHandlerReturnsResult) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -76,10 +107,11 @@ TEST_F(ServerToolTest, NonTemplateSyncHandlerReturnsResult) { ASSERT_EQ(responses.size(), 2); EXPECT_EQ(responses[1]["id"], "2"); ASSERT_TRUE(responses[1].contains("result")); - EXPECT_EQ(responses[1]["result"]["message"], "Hello, Alice!"); + EXPECT_EQ(responses[1]["result"]["structuredContent"]["message"], "Hello, Alice!"); + EXPECT_EQ(responses[1]["result"]["content"][0]["type"], "text"); } -TEST_F(ServerToolTest, NonTemplateHandlerExceptionBecomesJsonRpcError) { +TEST_F(ServerToolTest, NonTemplateHandlerExceptionBecomesToolError) { mcp::ServerCapabilities caps; caps.tools = mcp::ServerCapabilities::ToolsCapability{}; ServerSetup setup(io_ctx_, std::move(caps)); @@ -97,6 +129,7 @@ TEST_F(ServerToolTest, NonTemplateHandlerExceptionBecomesJsonRpcError) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json call_req; call_req["jsonrpc"] = "2.0"; @@ -124,6 +157,156 @@ TEST_F(ServerToolTest, NonTemplateHandlerExceptionBecomesJsonRpcError) { EXPECT_NE(err_text.find("something broke"), std::string::npos); } +TEST_F(ServerToolTest, NonStdExceptionFromHandlerBecomesToolError) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool("fail", "Throws something that is not a std::exception", + nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> nlohmann::json { throw 42; }); + + std::vector responses; + boost::asio::steady_timer watchdog(io_ctx_); + + setup.raw_transport->set_on_write([&responses, &setup, &watchdog](std::string_view message) { + responses.push_back(nlohmann::json::parse(message)); + if (responses.size() == 2) { + watchdog.cancel(); + setup.raw_transport->close(); + } + }); + + setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); + setup.raw_transport->enqueue_message( + nlohmann::json{{"jsonrpc", "2.0"}, + {"id", "2"}, + {"method", "tools/call"}, + {"params", {{"name", "fail"}, {"arguments", nlohmann::json::object()}}}} + .dump()); + + // A throw the tool guard does not catch escapes the dispatcher without writing anything, so + // the request is never answered and nothing closes the transport. Bound the wait, or that + // regression hangs the suite instead of failing it. + watchdog.expires_after(std::chrono::seconds(5)); + watchdog.async_wait([&setup](const boost::system::error_code& error) { + if (!error) { + setup.raw_transport->close(); + } + }); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await setup.server.run(setup.transport, io_ctx_.get_executor()); + }, + boost::asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(responses.size(), 2) << "the tools/call request went unanswered"; + EXPECT_EQ(responses[1]["id"], "2"); + ASSERT_TRUE(responses[1].contains("result")) << responses[1].dump(); + EXPECT_TRUE(responses[1]["result"]["isError"].get()); +} + +TEST_F(ServerToolTest, NonTemplateSyncHandlerReturningCallToolResultIsNotRewrapped) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool("greet", "Greets a user", nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> nlohmann::json { + return nlohmann::json(mcp::make_tool_text_result("Hello, Alice!")); + }); + + auto responses = call_tool(setup, "greet"); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_EQ(responses[1]["result"]["content"][0]["text"], "Hello, Alice!"); + EXPECT_FALSE(responses[1]["result"].contains("structuredContent")); +} + +TEST_F(ServerToolTest, NonTemplateSyncHandlerReturningErrorResultKeepsItsErrorFlag) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool("deny", "Denies every call", nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> nlohmann::json { + return nlohmann::json(mcp::make_tool_error_result("not permitted")); + }); + + auto responses = call_tool(setup, "deny"); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_TRUE(responses[1]["result"]["isError"].get()); + EXPECT_EQ(responses[1]["result"]["content"][0]["text"], "not permitted"); +} + +TEST_F(ServerToolTest, NonTemplateSyncHandlerReturningLiteralResultJsonIsNotRewrapped) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool( + "echo", "Echoes a message", nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> nlohmann::json { + return nlohmann::json{{"content", nlohmann::json::array({nlohmann::json{ + {"type", "text"}, {"text", "hello"}}})}}; + }); + + auto responses = call_tool(setup, "echo"); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("result")); + ASSERT_EQ(responses[1]["result"]["content"].size(), 1); + EXPECT_EQ(responses[1]["result"]["content"][0]["text"], "hello"); + EXPECT_FALSE(responses[1]["result"].contains("structuredContent")); +} + +TEST_F(ServerToolTest, DomainObjectNamedContentIsStillWrapped) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool( + "fetch", "Returns domain data", nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> nlohmann::json { + return nlohmann::json{{"content", nlohmann::json::array({"first", "second"})}}; + }); + + auto responses = call_tool(setup, "fetch"); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_EQ(responses[1]["result"]["structuredContent"]["content"][0], "first"); + EXPECT_EQ(responses[1]["result"]["content"][0]["type"], "text"); +} + +TEST_F(ServerToolTest, SyncHandlerExceptionTextIsSanitizedBeforeItReachesThePeer) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + setup.server.add_tool("fail", "Always fails", nlohmann::json{{"type", "object"}}, + [](const nlohmann::json&) -> nlohmann::json { + throw std::runtime_error("broke\nat /opt/internal/secret.cpp:12"); + }); + + auto responses = call_tool(setup, "fail"); + + ASSERT_EQ(responses.size(), 2); + ASSERT_TRUE(responses[1].contains("result")); + EXPECT_TRUE(responses[1]["result"]["isError"].get()); + auto text = responses[1]["result"]["content"][0]["text"].get(); + EXPECT_EQ(text.find('\n'), std::string::npos); + EXPECT_NE(text.find("broke"), std::string::npos); +} + TEST_F(ServerToolTest, NonTemplateToolAppearsInToolsList) { mcp::ServerCapabilities caps; caps.tools = mcp::ServerCapabilities::ToolsCapability{}; @@ -142,6 +325,7 @@ TEST_F(ServerToolTest, NonTemplateToolAppearsInToolsList) { }); setup.raw_transport->enqueue_message(make_initialize_request("1").dump()); + setup.raw_transport->enqueue_message(make_initialized_notification().dump()); nlohmann::json list_req; list_req["jsonrpc"] = "2.0"; @@ -166,3 +350,75 @@ TEST_F(ServerToolTest, NonTemplateToolAppearsInToolsList) { EXPECT_EQ(tools_arr[0]["description"], "Echoes input"); EXPECT_EQ(tools_arr[0]["inputSchema"], schema); } + +TEST_F(ServerToolTest, ResultHelpersProduceProtocolValidShapes) { + nlohmann::json text = mcp::make_tool_text_result("done"); + ASSERT_EQ(text["content"].size(), 1); + EXPECT_EQ(text["content"][0], nlohmann::json({{"type", "text"}, {"text", "done"}})); + + nlohmann::json error = mcp::make_tool_error_result("failed"); + EXPECT_TRUE(error["isError"].get()); + EXPECT_EQ(error["content"][0]["text"], "failed"); + + nlohmann::json structured = mcp::make_tool_structured_result({{"answer", 42}}, "forty-two"); + EXPECT_EQ(structured["structuredContent"]["answer"], 42); + EXPECT_EQ(structured["content"][0]["text"], "forty-two"); +} + +TEST_F(ServerToolTest, StructuredResultRoundTripsScalarContent) { + nlohmann::json structured = mcp::make_tool_structured_result(42); + EXPECT_EQ(structured["structuredContent"], 42); + EXPECT_EQ(structured["content"][0]["text"], "42"); +} + +TEST_F(ServerToolTest, StructuredResultRoundTripsArrayContent) { + nlohmann::json array_content = nlohmann::json::array({1, 2, 3}); + nlohmann::json structured = mcp::make_tool_structured_result(array_content); + EXPECT_EQ(structured["structuredContent"], array_content); + EXPECT_EQ(structured["content"][0]["text"], array_content.dump()); +} + +TEST_F(ServerToolTest, StructuredResultRoundTripsNullContent) { + nlohmann::json structured = mcp::make_tool_structured_result(nlohmann::json(nullptr)); + EXPECT_TRUE(structured["structuredContent"].is_null()); + EXPECT_EQ(structured["content"][0]["text"], "null"); +} + +TEST_F(ServerToolTest, RawToolAcceptsCompleteCallToolResult) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::Tool tool; + tool.name = "raw"; + tool.inputSchema = {{"type", "object"}}; + setup.server.add_raw_tool(tool, [](const nlohmann::json&) -> nlohmann::json { + return mcp::make_tool_text_result("raw result"); + }); + + auto responses = call_tool(setup, "raw"); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["result"]["content"][0]["text"], "raw result"); + EXPECT_FALSE(responses[1]["result"].contains("structuredContent")); +} + +TEST_F(ServerToolTest, RawToolRejectsIncompleteCallToolResult) { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + ServerSetup setup(io_ctx_, std::move(caps)); + + mcp::Tool tool; + tool.name = "invalid-raw"; + tool.inputSchema = {{"type", "object"}}; + setup.server.add_raw_tool(tool, [](const nlohmann::json&) -> nlohmann::json { + return {{"message", "not a CallToolResult"}}; + }); + + auto responses = call_tool(setup, "invalid-raw"); + + ASSERT_EQ(responses.size(), 2); + EXPECT_EQ(responses[1]["error"]["code"], mcp::g_INTERNAL_ERROR); + EXPECT_NE(responses[1]["error"]["message"].get().find("CallToolResult"), + std::string::npos); +} diff --git a/test/support/resolve_gate.cpp b/test/support/resolve_gate.cpp new file mode 100644 index 0000000..2e7862c --- /dev/null +++ b/test/support/resolve_gate.cpp @@ -0,0 +1,39 @@ +#include "resolve_gate.hpp" + +#ifdef __linux__ + +#include +#include + +ResolveGate& resolve_gate() { + static ResolveGate gate; + return gate; +} + +// A sanitizer runtime linked into the executable itself, which is how Clang links them, defines +// its interceptor as a weak getaddrinfo(). The definition below replaces it, and RTLD_NEXT then +// finds libc and skips the interceptor. The runtime exports the interceptor under this second name +// for that case; the reference is null when no such runtime is linked. MemorySanitizer depends on +// it: it learns that the returned list is initialised only by seeing the call. +extern "C" int __interceptor_getaddrinfo( // NOLINT(bugprone-reserved-identifier) + const char* node, const char* service, const struct addrinfo* hints, struct addrinfo** result) + __attribute__((weak)); + +// Interposes the libc symbol for this test binary only. A definition in the executable is found +// ahead of libc by the SDK's getaddrinfo() calls, whether the SDK is linked shared or static; +// everything is forwarded to a sanitizer's interceptor if there is one, and otherwise to the next +// definition in line. +extern "C" int getaddrinfo(const char* node, const char* service, const struct addrinfo* hints, + struct addrinfo** result) { + using GetAddrInfo = int (*)(const char*, const char*, const struct addrinfo*, struct addrinfo**); + static const auto next = __interceptor_getaddrinfo != nullptr + ? &__interceptor_getaddrinfo + : reinterpret_cast(dlsym(RTLD_NEXT, "getaddrinfo")); + if (next == nullptr) { + return EAI_FAIL; + } + resolve_gate().hold_if_armed(service); + return next(node, service, hints, result); +} + +#endif // __linux__ diff --git a/test/support/resolve_gate.hpp b/test/support/resolve_gate.hpp new file mode 100644 index 0000000..54c12da --- /dev/null +++ b/test/support/resolve_gate.hpp @@ -0,0 +1,64 @@ +#pragma once + +// Linux only: the gate is driven by a getaddrinfo() definition in the test binary (see +// resolve_gate.cpp), which is how a test parks a real Asio resolve past its cancel check. +#ifdef __linux__ + +#include +#include +#include +#include +#include + +/// Holds one getaddrinfo() call for an armed port inside the call, which is past the only point at +/// which Asio's resolver thread looks at its cancel token. Calls for any other port pass straight +/// through, so only the test that arms the gate is affected. Every wait is bounded. +class ResolveGate final { + public: + void arm(unsigned short port) { + std::lock_guard lock(mutex_); + entered_ = false; + released_ = false; + armed_port_.store(port); + } + + /// Let a held call go and stop holding new ones. + void release() { + armed_port_.store(0); + { + std::lock_guard lock(mutex_); + released_ = true; + } + changed_.notify_all(); + } + + [[nodiscard]] bool wait_until_entered(std::chrono::seconds limit) { + std::unique_lock lock(mutex_); + return changed_.wait_for(lock, limit, [this]() { return entered_; }); + } + + void hold_if_armed(const char* service) { + const auto port = armed_port_.load(); + if (port == 0 || service == nullptr || std::to_string(port) != service) { + return; + } + std::unique_lock lock(mutex_); + entered_ = true; + changed_.notify_all(); + // Longer than any wait in the test, and still bounded: a test that fails without + // releasing cannot leave the resolver thread parked forever. + changed_.wait_for(lock, std::chrono::seconds(30), [this]() { return released_; }); + } + + private: + std::mutex mutex_; + std::condition_variable changed_; + std::atomic armed_port_{0}; + bool entered_{false}; + bool released_{false}; +}; + +/// The one gate the binary's getaddrinfo() consults. +ResolveGate& resolve_gate(); + +#endif // __linux__ diff --git a/test/support/socket_gate.cpp b/test/support/socket_gate.cpp new file mode 100644 index 0000000..08614c1 --- /dev/null +++ b/test/support/socket_gate.cpp @@ -0,0 +1,31 @@ +#include "socket_gate.hpp" + +#ifdef __linux__ + +#include +#include + +SocketGate& socket_gate() { + static SocketGate gate; + return gate; +} + +// Interposed the same way as getaddrinfo() in resolve_gate.cpp, including the forward to a +// sanitizer's interceptor. noexcept because glibc declares it so. +extern "C" int __interceptor_socket( // NOLINT(bugprone-reserved-identifier) + int domain, int type, int protocol) noexcept __attribute__((weak)); + +extern "C" int socket(int domain, int type, int protocol) noexcept { + using Socket = int (*)(int, int, int) noexcept; + static const auto next = __interceptor_socket != nullptr + ? &__interceptor_socket + : reinterpret_cast(dlsym(RTLD_NEXT, "socket")); + if (next == nullptr) { + errno = ENOSYS; + return -1; + } + socket_gate().hold_if_armed(domain, type); + return next(domain, type, protocol); +} + +#endif // __linux__ diff --git a/test/support/socket_gate.hpp b/test/support/socket_gate.hpp new file mode 100644 index 0000000..db48291 --- /dev/null +++ b/test/support/socket_gate.hpp @@ -0,0 +1,74 @@ +#pragma once + +// Linux only: the gate is driven by a socket() definition in the test binary (see +// socket_gate.cpp), which is how a test parks an Asio connect at the point its socket is opened. +#ifdef __linux__ + +#include +#include +#include +#include +#include + +/// Holds one socket() call for a TCP socket inside the call, once armed. Inside async_connect() +/// that call is the socket being opened, so a held exchange has passed its last look at the abort +/// latch and has nothing open yet. Calls before and after the held one pass straight through. +/// Every wait is bounded. +class SocketGate final { + public: + /// Hold the TCP socket() call that follows `calls_to_pass` others. + void arm(int calls_to_pass) { + std::lock_guard lock(mutex_); + entered_ = false; + released_ = false; + countdown_.store(calls_to_pass + 1); + } + + /// Let a held call go and stop holding new ones. + void release() { + countdown_.store(0); + { + std::lock_guard lock(mutex_); + released_ = true; + } + changed_.notify_all(); + } + + [[nodiscard]] bool wait_until_entered(std::chrono::seconds limit) { + std::unique_lock lock(mutex_); + return changed_.wait_for(lock, limit, [this]() { return entered_; }); + } + + void hold_if_armed(int domain, int type) { + if (domain != AF_INET || (type & ~(SOCK_NONBLOCK | SOCK_CLOEXEC)) != SOCK_STREAM) { + return; + } + auto countdown = countdown_.load(); + do { + if (countdown == 0) { + return; + } + } while (!countdown_.compare_exchange_weak(countdown, countdown - 1)); + if (countdown != 1) { + return; + } + std::unique_lock lock(mutex_); + entered_ = true; + changed_.notify_all(); + // Longer than any wait in the test, and still bounded: a test that fails without + // releasing cannot leave an io thread parked forever. + changed_.wait_for(lock, std::chrono::seconds(30), [this]() { return released_; }); + } + + private: + std::mutex mutex_; + std::condition_variable changed_; + std::atomic countdown_{0}; + bool entered_{false}; + bool released_{false}; +}; + +/// The one gate the binary's socket() consults. +SocketGate& socket_gate(); + +#endif // __linux__ diff --git a/test/support/stalling_server.hpp b/test/support/stalling_server.hpp new file mode 100644 index 0000000..394b5cf --- /dev/null +++ b/test/support/stalling_server.hpp @@ -0,0 +1,39 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include + +/// Accepts one connection and holds it open without ever reading or writing on it: a peer that +/// stalls, which `close()` must be able to get away from without waiting for a timeout. +class StallingServer final { + public: + explicit StallingServer(boost::asio::io_context& io_ctx) + : acceptor_(io_ctx, {boost::asio::ip::make_address("127.0.0.1"), 0}) {} + + [[nodiscard]] unsigned short port() const { return acceptor_.local_endpoint().port(); } + + /// Start accepting; the accepted socket is held as a member so the connection stays open (no + /// FIN, no RST) until this server is destroyed. + void accept_and_stall() { + acceptor_.async_accept( + [this](boost::system::error_code error, boost::asio::ip::tcp::socket socket) { + if (!error) { + held_socket_ = std::move(socket); + accepted_.fetch_add(1); + } + }); + } + + /// Safe to poll from a thread other than the one running the io_context. + [[nodiscard]] int accepted() const { return accepted_.load(); } + + private: + boost::asio::ip::tcp::acceptor acceptor_; + std::optional held_socket_; + std::atomic accepted_{0}; +}; diff --git a/test/test_utils.hpp b/test/test_utils.hpp index b164a89..691dd7f 100644 --- a/test/test_utils.hpp +++ b/test/test_utils.hpp @@ -109,6 +109,10 @@ inline nlohmann::json make_shutdown_request(std::string_view id) { return {{"jsonrpc", "2.0"}, {"id", id}, {"method", "shutdown"}}; } +inline nlohmann::json make_initialized_notification() { + return {{"jsonrpc", "2.0"}, {"method", "notifications/initialized"}}; +} + inline nlohmann::json make_tool_call_request(std::string_view id, const std::string& tool_name, nlohmann::json arguments = nlohmann::json::object()) { return {{"jsonrpc", "2.0"}, diff --git a/test/transport/transport_http_session_manager_test.cpp b/test/transport/transport_http_session_manager_test.cpp index 765c0f4..c9d3810 100644 --- a/test/transport/transport_http_session_manager_test.cpp +++ b/test/transport/transport_http_session_manager_test.cpp @@ -1,16 +1,30 @@ +#include "mcp/auth/oauth.hpp" #include "mcp/transport/http_session_manager.hpp" #include +#include #include #include #include +#include +#include +#include +#include #include +#include #include #include +#include +#include +#include #include #include +#include #include +#include +#include +#include namespace { @@ -30,6 +44,7 @@ struct RawResponse { std::string session_id; std::string content_type; std::string allow; + std::string www_authenticate; }; /// Fire a single HTTP request and return the response. @@ -37,7 +52,9 @@ struct RawResponse { mcp::Task raw_request( const asio::any_io_executor& executor, unsigned short port, http::verb method, const std::string& target, const std::string& body = {}, const std::string& session_id = {}, - std::optional protocol_version = std::string(mcp::g_LATEST_PROTOCOL_VERSION)) { + std::optional protocol_version = std::string(mcp::g_LATEST_PROTOCOL_VERSION), + std::optional origin = std::nullopt, + std::optional bearer_token = std::nullopt, std::string accept = "application/json") { beast::tcp_stream stream(executor); auto resolver = asio::ip::tcp::resolver(executor); auto endpoints = @@ -47,7 +64,7 @@ mcp::Task raw_request( http::request request{method, target, 11}; request.set(http::field::host, "127.0.0.1"); request.set(http::field::content_type, "application/json"); - request.set(http::field::accept, "application/json"); + request.set(http::field::accept, accept); if (protocol_version.has_value()) { request.set("MCP-Protocol-Version", *protocol_version); @@ -56,6 +73,12 @@ mcp::Task raw_request( if (!session_id.empty()) { request.set("Mcp-Session-Id", session_id); } + if (origin.has_value()) { + request.set(http::field::origin, *origin); + } + if (bearer_token.has_value()) { + request.set(http::field::authorization, "Bearer " + *bearer_token); + } if (!body.empty()) { request.body() = body; @@ -73,6 +96,7 @@ mcp::Task raw_request( result.body = response.body(); result.content_type = std::string(response[http::field::content_type]); result.allow = std::string(response[http::field::allow]); + result.www_authenticate = std::string(response[http::field::www_authenticate]); auto session_it = response.find("Mcp-Session-Id"); if (session_it != response.end()) { @@ -85,6 +109,24 @@ mcp::Task raw_request( co_return result; } +std::vector> parse_sse_events(const std::string& body) { + std::vector> events; + std::istringstream stream(body); + std::string event_id; + std::string line; + + while (std::getline(stream, line)) { + if (line.starts_with("id: ")) { + event_id = line.substr(4); + } else if (line.starts_with("data: ")) { + events.emplace_back(event_id, json::parse(line.substr(6))); + event_id.clear(); + } + } + + return events; +} + // --------------------------------------------------------------------------- // Helper: create a ServerFactory for the session manager // --------------------------------------------------------------------------- @@ -96,15 +138,31 @@ mcp::StreamableHttpSessionManager::ServerFactory make_echo_server_factory() { auto server = std::make_unique(mcp::Implementation{"test-session-server", "1.0.0"}, std::move(caps)); - server->add_tool( + server->add_tool( "echo", "Echoes input", json{{"type", "object"}, {"properties", {{"message", {{"type", "string"}}}}}}, - [](json params) -> mcp::Task { - json result; - result["content"] = json::array(); - result["content"].push_back( - json{{"type", "text"}, {"text", params.value("message", "empty")}}); - co_return result; + [](json params) -> mcp::Task { + co_return mcp::make_tool_text_result(params.value("message", "empty")); + }); + + return server; + }; +} + +mcp::StreamableHttpSessionManager::ServerFactory make_notifying_server_factory() { + return [](const asio::any_io_executor&) -> std::unique_ptr { + mcp::ServerCapabilities caps; + caps.tools = mcp::ServerCapabilities::ToolsCapability{}; + caps.logging = json::object(); + auto server = std::make_unique( + mcp::Implementation{"test-notifying-server", "1.0.0"}, std::move(caps)); + + server->add_tool( + "notify", "Sends notifications before returning", json{{"type", "object"}}, + [](mcp::Context& context, json) -> mcp::Task { + co_await context.log_info("tool started"); + co_await context.report_progress(0.5, 1.0, "halfway"); + co_return mcp::make_tool_text_result("done"); }); return server; @@ -128,6 +186,15 @@ mcp::Task do_initialize( std::move(header_protocol_version)); } +mcp::Task send_initialized( + const asio::any_io_executor& executor, unsigned short port, const std::string& session_id, + std::optional header_protocol_version = std::string(mcp::g_LATEST_PROTOCOL_VERSION)) { + json notification = {{"jsonrpc", "2.0"}, {"method", "notifications/initialized"}}; + + co_return co_await raw_request(executor, port, http::verb::post, "/mcp", notification.dump(), + session_id, std::move(header_protocol_version)); +} + } // namespace // =========================================================================== @@ -139,6 +206,379 @@ class SessionManagerTest : public ::testing::Test { asio::io_context io_ctx_; }; +TEST_F(SessionManagerTest, SsePostReturnsNotificationsThenFinalResponseInOrder) { + const unsigned short port = 19111; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_notifying_server_factory()); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse tool_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto init = co_await do_initialize(io_ctx_.get_executor(), port); + EXPECT_EQ(init.status, 200); + auto initialized = co_await send_initialized(io_ctx_.get_executor(), port, init.session_id); + EXPECT_EQ(initialized.status, 202); + + json request = {{"jsonrpc", "2.0"}, + {"method", "tools/call"}, + {"params", + {{"name", "notify"}, + {"arguments", json::object()}, + {"_meta", {{"progressToken", "progress-1"}}}}}, + {"id", 2}}; + tool_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", request.dump(), init.session_id, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), std::nullopt, std::nullopt, + "application/json, text/event-stream"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(tool_response.status, 200); + EXPECT_EQ(tool_response.content_type, "text/event-stream"); + + const auto events = parse_sse_events(tool_response.body); + ASSERT_EQ(events.size(), 3); + EXPECT_EQ(events[0].second.at("method"), "notifications/message"); + EXPECT_EQ(events[0].second.at("params").at("data"), "tool started"); + EXPECT_EQ(events[1].second.at("method"), "notifications/progress"); + EXPECT_EQ(events[1].second.at("params").at("progressToken"), "progress-1"); + EXPECT_EQ(events[2].second.at("id"), 2); + EXPECT_EQ(events[2].second.at("result").at("content").at(0).at("text"), "done"); + EXPECT_LT(std::stoull(events[0].first), std::stoull(events[1].first)); + EXPECT_LT(std::stoull(events[1].first), std::stoull(events[2].first)); +} + +TEST_F(SessionManagerTest, JsonOnlyPostSkipsNotificationsAndReturnsFinalResponse) { + const unsigned short port = 19112; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_notifying_server_factory()); + manager.set_json_only(true); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse tool_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto init = co_await do_initialize(io_ctx_.get_executor(), port); + EXPECT_EQ(init.status, 200); + auto initialized = co_await send_initialized(io_ctx_.get_executor(), port, init.session_id); + EXPECT_EQ(initialized.status, 202); + + json request = {{"jsonrpc", "2.0"}, + {"method", "tools/call"}, + {"params", + {{"name", "notify"}, + {"arguments", json::object()}, + {"_meta", {{"progressToken", "progress-1"}}}}}, + {"id", 2}}; + tool_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", request.dump(), init.session_id, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), std::nullopt, std::nullopt, + "application/json, text/event-stream"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(tool_response.status, 200); + EXPECT_EQ(tool_response.content_type, "application/json"); + const auto response = json::parse(tool_response.body); + EXPECT_EQ(response.at("id"), 2); + EXPECT_EQ(response.at("result").at("content").at(0).at("text"), "done"); +} + +TEST_F(SessionManagerTest, OriginHeadersAreDeniedUntilExplicitlyAllowed) { + const unsigned short port = 19094; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_allowed_origins({"https://trusted.example"}); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse denied_response; + RawResponse allowed_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto init_body = + json{{"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"clientInfo", {{"name", "test"}, {"version", "1"}}}, + {"capabilities", json::object()}}}, + {"id", 1}} + .dump(); + denied_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", init_body, {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), "https://untrusted.example"); + allowed_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", init_body, {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), "https://trusted.example"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(denied_response.status, 403); + EXPECT_EQ(allowed_response.status, 200); +} + +namespace { + +/// The initialize body the origin tests below post; none of them gets far enough to care what is +/// in it, only whether the origin check let it through. +std::string make_initialize_body() { + return json{{"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"clientInfo", {{"name", "test"}, {"version", "1"}}}, + {"capabilities", json::object()}}}, + {"id", 1}} + .dump(); +} + +} // namespace + +// The allow-list is a set of exact strings, not the canonicalizing comparison the client-side +// MetadataFetchPolicy performs on the origins it will fetch from. An explicit default port, a +// different scheme case and a trailing slash all denote the same web origin, and all three are +// refused here. A deployment that wants them admitted must list every spelling it will receive. +TEST_F(SessionManagerTest, AllowedOriginComparisonIsExactRatherThanCanonicalizing) { + const unsigned short port = 19150; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_allowed_origins({"https://trusted.example"}); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + const std::vector spellings = {"https://trusted.example:443", + "HTTPS://trusted.example", "https://Trusted.example", + "https://trusted.example/"}; + std::vector statuses; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + for (const auto& origin : spellings) { + const auto response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", make_initialize_body(), {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), origin); + statuses.push_back(response.status); + } + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(statuses.size(), spellings.size()); + for (std::size_t index = 0; index < spellings.size(); ++index) { + SCOPED_TRACE(spellings[index]); + EXPECT_EQ(statuses[index], 403u); + } +} + +// set_allow_all_origins(true) is the documented escape hatch for a deployment that fronts the +// transport with its own origin policy. Nothing else has to be configured for it to take effect. +TEST_F(SessionManagerTest, AllowAllOriginsAdmitsAnOriginThatWasNeverListed) { + const unsigned short port = 19151; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_allow_all_origins(true); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", make_initialize_body(), {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), "https://untrusted.example"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(response.status, 200); +} + +// set_allowed_origins() names the origins that may connect, so it also revokes a blanket allowance +// granted earlier. Leaving allow-all in force behind an allow-list would make the narrower call +// silently do nothing. +TEST_F(SessionManagerTest, NamingAllowedOriginsRevokesAnEarlierAllowAll) { + const unsigned short port = 19152; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_allow_all_origins(true); + manager.set_allowed_origins({"https://trusted.example"}); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse denied_response; + RawResponse allowed_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", make_initialize_body(), {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), "https://untrusted.example"); + allowed_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", make_initialize_body(), {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), "https://trusted.example"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(denied_response.status, 403); + EXPECT_EQ(allowed_response.status, 200); +} + +// On this transport the origin check runs once, before everything else a request could be +// refused for, so a rebinding attempt is answered 403 without disclosing whether a token would +// have been accepted or what the protected-resource metadata says. HttpServerTransport reaches +// the same two answers by running the check per route instead, and the two diverge only on an +// unauthenticated path -- see the companion test in transport_http_test.cpp. +TEST_F(SessionManagerTest, TheOriginCheckPrecedesTheBearerCheckAndTheMetadataRoute) { + const unsigned short port = 19153; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_allowed_origins({"https://trusted.example"}); + manager.set_bearer_token_validator([](std::string_view token) { return token == "valid-token"; }); + manager.set_unauthenticated_paths({"/health"}); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "https://mcp.example.com/mcp"; + manager.set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse untokened_response; + RawResponse metadata_response; + RawResponse exempt_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + untokened_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", make_initialize_body(), {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), "https://untrusted.example"); + metadata_response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, + "/.well-known/oauth-protected-resource/mcp", {}, + {}, std::string(mcp::g_LATEST_PROTOCOL_VERSION), + "https://untrusted.example"); + exempt_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::get, "/health", {}, {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), "https://untrusted.example"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(untokened_response.status, 403); + EXPECT_TRUE(untokened_response.www_authenticate.empty()); + EXPECT_EQ(metadata_response.status, 403); + + // HttpServerTransport answers this same request 404: it decides an exempt path is not served + // before any origin check runs. Neither status discloses anything a disallowed origin could + // not already infer, but the two transports do not agree on which one to send. + EXPECT_EQ(exempt_response.status, 403); +} + +TEST_F(SessionManagerTest, BearerTokensAreValidatedAtHttpBoundary) { + const unsigned short port = 19095; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "valid-token"; }); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse missing_response; + RawResponse invalid_response; + RawResponse valid_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto init_body = + json{{"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"clientInfo", {{"name", "test"}, {"version", "1"}}}, + {"capabilities", json::object()}}}, + {"id", 1}} + .dump(); + missing_response = + co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, "/mcp", init_body); + invalid_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", init_body, {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), std::nullopt, "invalid-token"); + valid_response = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", init_body, {}, + std::string(mcp::g_LATEST_PROTOCOL_VERSION), std::nullopt, "valid-token"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(missing_response.status, 401); + EXPECT_EQ(invalid_response.status, 401); + EXPECT_EQ(valid_response.status, 200); +} + +TEST_F(SessionManagerTest, SecurityChecksRunBeforeCustomHttpHandlers) { + const unsigned short port = 19110; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "valid-token"; }); + manager.set_custom_request_handler( + [](const mcp::StringRequest& request) -> std::optional { + if (request.target() != "/health") { + return std::nullopt; + } + mcp::StringResponse response{http::status::ok, request.version()}; + response.body() = "healthy"; + response.prepare_payload(); + return response; + }); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse missing_response; + RawResponse valid_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + missing_response = + co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/health"); + valid_response = + co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/health", {}, {}, + std::nullopt, std::nullopt, "valid-token"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(missing_response.status, 401); + EXPECT_EQ(valid_response.status, 200); + EXPECT_EQ(valid_response.body, "healthy"); +} + // --------------------------------------------------------------------------- // 1. Create session on initialize (no session header) // --------------------------------------------------------------------------- @@ -208,6 +648,9 @@ TEST_F(SessionManagerTest, SubsequentRequestsUseNegotiatedProtocolVersion) { [&]() -> mcp::Task { init_response = co_await do_initialize(io_ctx_.get_executor(), port, 1, "2025-06-18", std::nullopt); + auto initialized = co_await send_initialized(io_ctx_.get_executor(), port, + init_response.session_id, "2025-06-18"); + EXPECT_EQ(initialized.status, 202); json request = {{"jsonrpc", "2.0"}, {"method", "tools/call"}, @@ -226,8 +669,85 @@ TEST_F(SessionManagerTest, SubsequentRequestsUseNegotiatedProtocolVersion) { EXPECT_EQ(tool_response.status, 200); } +TEST_F(SessionManagerTest, ProtocolVersionStateIsSafeAcrossConnectionStrands) { + constexpr unsigned short port = 19115; + constexpr std::size_t request_count = 32; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + std::vector responses(request_count); + std::vector errors(request_count); + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto init = co_await do_initialize(io_ctx_.get_executor(), port); + EXPECT_EQ(init.status, 200); + + auto remaining = std::make_shared>(request_count); + auto all_done = std::make_shared(io_ctx_); + all_done->expires_at(std::chrono::steady_clock::time_point::max()); + + for (std::size_t index = 0; index < request_count; ++index) { + asio::co_spawn( + io_ctx_, + [&, index, session_id = init.session_id]() -> mcp::Task { + if (index % 2 == 0) { + json reinitialize = { + {"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", "2025-06-18"}, + {"clientInfo", {{"name", "race-test"}, {"version", "1"}}}, + {"capabilities", json::object()}}}, + {"id", static_cast(100 + index)}}; + responses[index] = co_await raw_request( + io_ctx_.get_executor(), port, http::verb::post, "/mcp", + reinitialize.dump(), session_id, "2025-06-18"); + } else { + const auto version = index % 4 == 1 + ? std::string("2025-06-18") + : std::string(mcp::g_LATEST_PROTOCOL_VERSION); + responses[index] = + co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, + "/mcp", {}, session_id, version); + } + }, + [&, index, remaining, all_done](std::exception_ptr error) { + errors[index] = std::move(error); + if (remaining->fetch_sub(1, std::memory_order_acq_rel) == 1) { + asio::post(all_done->get_executor(), [all_done]() { + all_done->expires_at(std::chrono::steady_clock::now()); + }); + } + }); + } + + boost::system::error_code wait_error; + co_await all_done->async_wait(asio::redirect_error(asio::use_awaitable, wait_error)); + manager.close(); + }, + asio::detached); + + std::vector workers; + workers.reserve(4); + for (int index = 0; index < 4; ++index) { + workers.emplace_back([this]() { io_ctx_.run(); }); + } + for (auto& worker : workers) { + worker.join(); + } + + for (std::size_t index = 0; index < request_count; ++index) { + EXPECT_EQ(errors[index], nullptr) << "request " << index; + EXPECT_TRUE(responses[index].status == 200 || responses[index].status == 400) + << "request " << index << " returned " << responses[index].status; + } +} + TEST_F(SessionManagerTest, SubsequentRequestsAllowMissingNegotiatedProtocolHeader) { - const unsigned short port = 19094; + const unsigned short port = 19113; mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, make_echo_server_factory()); @@ -241,6 +761,9 @@ TEST_F(SessionManagerTest, SubsequentRequestsAllowMissingNegotiatedProtocolHeade [&]() -> mcp::Task { init_response = co_await do_initialize(io_ctx_.get_executor(), port, 1, "2025-06-18", std::nullopt); + auto initialized = co_await send_initialized(io_ctx_.get_executor(), port, + init_response.session_id, std::nullopt); + EXPECT_EQ(initialized.status, 202); json request = {{"jsonrpc", "2.0"}, {"method", "tools/call"}, @@ -280,6 +803,8 @@ TEST_F(SessionManagerTest, RouteBySessionId) { auto init = co_await do_initialize(io_ctx_.get_executor(), port); EXPECT_EQ(init.status, 200); auto session_id = init.session_id; + auto initialized = co_await send_initialized(io_ctx_.get_executor(), port, session_id); + EXPECT_EQ(initialized.status, 202); // Call a tool using the session ID json tool_call = {{"jsonrpc", "2.0"}, @@ -388,10 +913,14 @@ TEST_F(SessionManagerTest, MultipleConcurrentSessionsSameIds) { // Create session A auto init_a = co_await do_initialize(io_ctx_.get_executor(), port, 1); auto session_a = init_a.session_id; + auto initialized_a = co_await send_initialized(io_ctx_.get_executor(), port, session_a); + EXPECT_EQ(initialized_a.status, 202); // Create session B auto init_b = co_await do_initialize(io_ctx_.get_executor(), port, 1); auto session_b = init_b.session_id; + auto initialized_b = co_await send_initialized(io_ctx_.get_executor(), port, session_b); + EXPECT_EQ(initialized_b.status, 202); EXPECT_NE(session_a, session_b); EXPECT_EQ(manager.session_count(), 2); @@ -583,77 +1112,103 @@ TEST_F(SessionManagerTest, NonInitializeWithoutSessionReturns400) { } // --------------------------------------------------------------------------- -// 10. DELETE on unknown session returns 404 +// server/discover is reachable with zero prior state in stateful mode: no session is +// created and no Mcp-Session-Id is issued, unlike every other sessionless non-initialize +// method (pinned above by NonInitializeWithoutSessionReturns400). // --------------------------------------------------------------------------- -TEST_F(SessionManagerTest, DeleteUnknownSessionReturns404) { - const unsigned short port = 19089; +TEST_F(SessionManagerTest, SessionlessDiscoverReturns200WithoutCreatingSession) { + const unsigned short port = 19118; mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, make_echo_server_factory()); asio::co_spawn(io_ctx_, manager.listen(), asio::detached); RawResponse response; + std::size_t session_count = 0; asio::co_spawn( io_ctx_, [&]() -> mcp::Task { - response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::delete_, "/mcp", - {}, "nonexistent-session"); + json request = {{"jsonrpc", "2.0"}, {"id", 1}, {"method", "server/discover"}}; + response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + request.dump()); // no session_id + session_count = manager.session_count(); manager.close(); }, asio::detached); io_ctx_.run(); - EXPECT_EQ(response.status, 404); -} + EXPECT_EQ(response.status, 200); + EXPECT_TRUE(response.session_id.empty()); + EXPECT_EQ(session_count, 0); -// --------------------------------------------------------------------------- -// 11. Session count tracks sessions correctly -// --------------------------------------------------------------------------- + auto body = json::parse(response.body); + ASSERT_TRUE(body.contains("result")); + const auto& result = body["result"]; + EXPECT_EQ(result["resultType"], "complete"); + EXPECT_FALSE(result["supportedVersions"].get>().empty()); + ASSERT_TRUE(result.contains("_meta")); + EXPECT_EQ(result["_meta"]["io.modelcontextprotocol/serverInfo"]["name"], "test-session-server"); + ASSERT_TRUE(result["capabilities"].contains("tools")); + ASSERT_TRUE(result.contains("ttlMs")); + ASSERT_TRUE(result.contains("cacheScope")); +} -TEST_F(SessionManagerTest, SessionCountTracksCorrectly) { - const unsigned short port = 19090; +// This transport's immunity to the HttpServerTransport defect (sessionless discover rejected +// once a session exists) is NOT structural. It holds only because handle_post intercepts +// sessionless discover BEFORE calling resolve_session_for_post, which is two adjacent +// statements in a particular order. Any refactor that hoists session resolution earlier +// reintroduces that bug silently, so pin the after-a-session case, not just the zero-state +// case covered above. +TEST_F(SessionManagerTest, SessionlessDiscoverStaysReachableAfterSessionEstablished) { + const unsigned short port = 19121; mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, make_echo_server_factory()); asio::co_spawn(io_ctx_, manager.listen(), asio::detached); - std::size_t count_after_create = 0; - std::size_t count_after_second = 0; - std::size_t count_after_delete = 0; - + RawResponse init_response; + RawResponse discover_response; + std::size_t session_count_after = 0; asio::co_spawn( io_ctx_, [&]() -> mcp::Task { - auto init_a = co_await do_initialize(io_ctx_.get_executor(), port, 1); - count_after_create = manager.session_count(); - - auto init_b = co_await do_initialize(io_ctx_.get_executor(), port, 1); - count_after_second = manager.session_count(); - - // Delete the first session - co_await raw_request(io_ctx_.get_executor(), port, http::verb::delete_, "/mcp", {}, - init_a.session_id); - count_after_delete = manager.session_count(); + init_response = co_await do_initialize(io_ctx_.get_executor(), port); + json request = {{"jsonrpc", "2.0"}, {"id", 2}, {"method", "server/discover"}}; + discover_response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, + "/mcp", request.dump()); // no session id + session_count_after = manager.session_count(); manager.close(); }, asio::detached); io_ctx_.run(); - EXPECT_EQ(count_after_create, 1); - EXPECT_EQ(count_after_second, 2); - EXPECT_EQ(count_after_delete, 1); + ASSERT_EQ(init_response.status, 200); + ASSERT_FALSE(init_response.session_id.empty()); + + EXPECT_EQ(discover_response.status, 200); + EXPECT_TRUE(discover_response.session_id.empty()); + // The established session is untouched: still exactly one, and no second one was minted. + EXPECT_EQ(session_count_after, 1); + + auto body = json::parse(discover_response.body); + ASSERT_TRUE(body.contains("result")); + EXPECT_EQ(body["result"]["resultType"], "complete"); } // --------------------------------------------------------------------------- -// 12. Method not allowed returns 405 +// server/discover is exempted from the MCP-Protocol-Version header check that rejects any +// version outside g_SUPPORTED_PROTOCOL_VERSIONS, so a modern client probing with its own +// 2026-07-28 header is not rejected pre-dispatch — unlike every other method (pinned by +// BadProtocolVersionReturns400 for initialize, and StatelessBadProtocolVersionReturns400 / +// the discover-vs-tools-list test below for the general case). // --------------------------------------------------------------------------- -TEST_F(SessionManagerTest, UnsupportedMethodReturns405) { - const unsigned short port = 19091; +TEST_F(SessionManagerTest, SessionlessDiscoverAcceptsUnsupportedProtocolVersionHeader) { + const unsigned short port = 19119; mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, make_echo_server_factory()); @@ -663,19 +1218,158 @@ TEST_F(SessionManagerTest, UnsupportedMethodReturns405) { asio::co_spawn( io_ctx_, [&]() -> mcp::Task { - response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::put, "/mcp", - R"({"test":"data"})"); + json request = {{"jsonrpc", "2.0"}, {"id", 1}, {"method", "server/discover"}}; + response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + request.dump(), {}, "2026-07-28"); manager.close(); }, asio::detached); io_ctx_.run(); - EXPECT_EQ(response.status, 405); + EXPECT_EQ(response.status, 200); + auto body = json::parse(response.body); + ASSERT_TRUE(body.contains("result")); + EXPECT_EQ(body["result"]["resultType"], "complete"); +} + +// --------------------------------------------------------------------------- +// Stateless mode: server/discover with a 2026-07-28 header is accepted, while tools/list +// (non-discover) with the same header still gets the legacy 400 — pinning that the header +// exemption is scoped to server/discover and not widened to other methods. +// --------------------------------------------------------------------------- + +TEST_F(SessionManagerTest, + StatelessDiscoverAcceptsUnsupportedProtocolVersionHeaderToolsListStillRejected) { + const unsigned short port = 19120; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_stateless_json_mode(true); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse discover_response; + RawResponse tools_list_response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + json discover_request = {{"jsonrpc", "2.0"}, {"id", 1}, {"method", "server/discover"}}; + discover_response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, + "/mcp", discover_request.dump(), {}, "2026-07-28"); + + json tools_list_request = { + {"jsonrpc", "2.0"}, {"method", "tools/list"}, {"params", json::object()}, {"id", 2}}; + tools_list_response = + co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + tools_list_request.dump(), {}, "2026-07-28"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(discover_response.status, 200); + auto discover_body = json::parse(discover_response.body); + ASSERT_TRUE(discover_body.contains("result")); + EXPECT_EQ(discover_body["result"]["resultType"], "complete"); + + EXPECT_EQ(tools_list_response.status, 400); +} + +// --------------------------------------------------------------------------- +// 10. DELETE on unknown session returns 404 +// --------------------------------------------------------------------------- + +TEST_F(SessionManagerTest, DeleteUnknownSessionReturns404) { + const unsigned short port = 19089; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::delete_, "/mcp", + {}, "nonexistent-session"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(response.status, 404); +} + +// --------------------------------------------------------------------------- +// 11. Session count tracks sessions correctly +// --------------------------------------------------------------------------- + +TEST_F(SessionManagerTest, SessionCountTracksCorrectly) { + const unsigned short port = 19090; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + std::size_t count_after_create = 0; + std::size_t count_after_second = 0; + std::size_t count_after_delete = 0; + + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto init_a = co_await do_initialize(io_ctx_.get_executor(), port, 1); + count_after_create = manager.session_count(); + + auto init_b = co_await do_initialize(io_ctx_.get_executor(), port, 1); + count_after_second = manager.session_count(); + + // Delete the first session + co_await raw_request(io_ctx_.get_executor(), port, http::verb::delete_, "/mcp", {}, + init_a.session_id); + count_after_delete = manager.session_count(); + + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(count_after_create, 1); + EXPECT_EQ(count_after_second, 2); + EXPECT_EQ(count_after_delete, 1); +} + +// --------------------------------------------------------------------------- +// 12. Method not allowed returns 405 +// --------------------------------------------------------------------------- + +TEST_F(SessionManagerTest, UnsupportedMethodReturns405) { + const unsigned short port = 19091; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::put, "/mcp", + R"({"test":"data"})"); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(response.status, 405); } TEST_F(SessionManagerTest, StatelessInitializeDoesNotCreateSessionHeader) { - const unsigned short port = 19095; + const unsigned short port = 19114; mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, make_echo_server_factory()); manager.set_stateless_json_mode(true); @@ -822,3 +1516,645 @@ TEST_F(SessionManagerTest, StatelessBadProtocolVersionReturns400) { EXPECT_EQ(response.status, 400); EXPECT_TRUE(response.session_id.empty()); } + +TEST_F(SessionManagerTest, CloseCancelsIdleKeepAliveAndCannotCreatePostCloseSession) { + const unsigned short port = 19115; + std::atomic factory_calls{0}; + std::atomic custom_handler_calls{0}; + auto echo_factory = make_echo_server_factory(); + mcp::StreamableHttpSessionManager manager( + io_ctx_.get_executor(), "127.0.0.1", port, + [&factory_calls, echo_factory](const asio::any_io_executor& executor) mutable { + factory_calls.fetch_add(1, std::memory_order_relaxed); + return echo_factory(executor); + }); + manager.set_custom_request_handler([&custom_handler_calls](const mcp::StringRequest& request) + -> std::optional { + custom_handler_calls.fetch_add(1, std::memory_order_relaxed); + if (request.target() != "/health") { + return std::nullopt; + } + mcp::StringResponse response{http::status::ok, request.version()}; + response.set(http::field::content_type, "application/json"); + response.keep_alive(request.keep_alive()); + response.body() = R"({"status":"ok"})"; + response.prepare_payload(); + return response; + }); + + auto client = std::make_shared(io_ctx_.get_executor()); + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(2)); + bool first_request_completed = false; + bool connection_cancelled = false; + bool post_close_response_received = false; + bool timed_out = false; + std::exception_ptr client_error; + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + asio::co_spawn( + io_ctx_, + [&manager, client, deadline, port, &first_request_completed, &connection_cancelled, + &post_close_response_received]() -> mcp::Task { + asio::ip::tcp::resolver resolver(client->get_executor()); + auto endpoints = + co_await resolver.async_resolve("127.0.0.1", std::to_string(port), asio::use_awaitable); + co_await client->async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request health_request{http::verb::get, "/health", 11}; + health_request.set(http::field::host, "127.0.0.1"); + health_request.keep_alive(true); + co_await http::async_write(*client, health_request, asio::use_awaitable); + + beast::flat_buffer response_buffer; + http::response health_response; + co_await http::async_read(*client, response_buffer, health_response, asio::use_awaitable); + first_request_completed = health_response.result() == http::status::ok; + + manager.close(); + + http::request initialize_request{http::verb::post, "/mcp", 11}; + initialize_request.set(http::field::host, "127.0.0.1"); + initialize_request.set(http::field::content_type, "application/json"); + initialize_request.set("MCP-Protocol-Version", mcp::g_LATEST_PROTOCOL_VERSION); + initialize_request.keep_alive(true); + initialize_request.body() = + R"({"jsonrpc":"2.0","method":"initialize","params":{"protocolVersion":"2025-06-18","clientInfo":{"name":"post-close","version":"1.0"},"capabilities":{}},"id":1})"; + initialize_request.prepare_payload(); + + try { + co_await http::async_write(*client, initialize_request, asio::use_awaitable); + http::response initialize_response; + co_await http::async_read(*client, response_buffer, initialize_response, + asio::use_awaitable); + post_close_response_received = true; + } catch (const boost::system::system_error&) { + connection_cancelled = true; + } + deadline->cancel(); + }, + [&client_error, deadline, &manager](std::exception_ptr error) { + client_error = std::move(error); + deadline->cancel(); + manager.close(); + }); + deadline->async_wait([client, &manager, &timed_out](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + manager.close(); + boost::system::error_code ignored; + (void)client->socket().close(ignored); + }); + + io_ctx_.run(); + + EXPECT_EQ(client_error, nullptr); + EXPECT_FALSE(timed_out); + EXPECT_TRUE(first_request_completed); + EXPECT_TRUE(connection_cancelled); + EXPECT_FALSE(post_close_response_received); + EXPECT_EQ(custom_handler_calls.load(std::memory_order_relaxed), 1); + EXPECT_EQ(factory_calls.load(std::memory_order_relaxed), 0); + EXPECT_EQ(manager.session_count(), 0); +} + +TEST_F(SessionManagerTest, NonAtomicConfigurationLocksWhenListeningStarts) { + const unsigned short port = 19116; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + auto listener = manager.listen(); + + EXPECT_NO_THROW(manager.set_json_only(true)); + EXPECT_THROW(manager.set_custom_request_handler({}), std::logic_error); + EXPECT_THROW(manager.set_allowed_origins({"https://trusted.example"}), std::logic_error); + EXPECT_THROW(manager.set_allow_all_origins(true), std::logic_error); + EXPECT_THROW(manager.set_bearer_token_validator({}), std::logic_error); + EXPECT_THROW(manager.set_stateless_json_mode(true), std::logic_error); + EXPECT_THROW(manager.set_tool_executor(io_ctx_.get_executor()), std::logic_error); + EXPECT_THROW(manager.set_bearer_challenge({}), std::logic_error); + EXPECT_THROW(manager.set_protected_resource_metadata({}), std::logic_error); + EXPECT_THROW(manager.set_unauthenticated_paths({"/health"}), std::logic_error); + EXPECT_THROW(manager.set_async_bearer_token_validator({}), std::logic_error); + EXPECT_THROW(manager.set_max_request_body_bytes(4096), std::logic_error); + + manager.close(); + asio::co_spawn(io_ctx_, std::move(listener), asio::detached); + io_ctx_.run(); +} + +TEST_F(SessionManagerTest, StatelessDispatchUsesConfiguredToolExecutor) { + const unsigned short port = 19117; + const auto http_thread = std::this_thread::get_id(); + std::atomic handler_ran{false}; + std::atomic handler_ran_on_http_thread{true}; + asio::thread_pool tool_pool(1); + + mcp::StreamableHttpSessionManager manager( + io_ctx_.get_executor(), "127.0.0.1", port, + [&handler_ran, &handler_ran_on_http_thread, + http_thread](const asio::any_io_executor&) -> std::unique_ptr { + mcp::ServerCapabilities capabilities; + capabilities.tools = mcp::ServerCapabilities::ToolsCapability{}; + auto server = std::make_unique( + mcp::Implementation{"executor-test-server", "1.0.0"}, std::move(capabilities)); + server->add_tool( + "executor", "Reports executor affinity", json{{"type", "object"}}, + [&handler_ran, &handler_ran_on_http_thread, + http_thread](json) -> mcp::Task { + handler_ran_on_http_thread.store(std::this_thread::get_id() == http_thread, + std::memory_order_release); + handler_ran.store(true, std::memory_order_release); + co_return mcp::make_tool_text_result("tool executor"); + }); + return server; + }); + manager.set_stateless_json_mode(true); + manager.set_tool_executor(tool_pool.get_executor()); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + RawResponse response; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const json request = {{"jsonrpc", "2.0"}, + {"method", "tools/call"}, + {"params", {{"name", "executor"}, {"arguments", json::object()}}}, + {"id", 1}}; + response = co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + request.dump()); + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + tool_pool.join(); + + EXPECT_EQ(response.status, 200); + EXPECT_TRUE(handler_ran.load(std::memory_order_acquire)); + EXPECT_FALSE(handler_ran_on_http_thread.load(std::memory_order_acquire)); + EXPECT_EQ(json::parse(response.body)["result"]["content"][0]["text"], "tool executor"); +} + +// =========================================================================== +// Server-side OAuth challenge, metadata route and unauthenticated paths +// =========================================================================== + +TEST_F(SessionManagerTest, UnconfiguredManagerSendsBareBearerChallenge) { + const unsigned short port = 19130; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/mcp"); + manager.close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(denied.www_authenticate, "Bearer"); +} + +TEST_F(SessionManagerTest, ConfiguredChallengeIsSentOnUnauthorized) { + const unsigned short port = 19131; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + + mcp::BearerChallengeConfig challenge; + challenge.realm = "mcp"; + challenge.error = "invalid_token"; + challenge.scope = "mcp:read"; + challenge.resource_metadata = "http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp"; + manager.set_bearer_challenge(challenge); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/mcp"); + manager.close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(denied.www_authenticate, + R"(Bearer realm="mcp", error="invalid_token", scope="mcp:read", )" + R"(resource_metadata="http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp")"); +} + +TEST_F(SessionManagerTest, ProtectedResourceMetadataIsReadableWithoutAToken) { + const unsigned short port = 19132; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "http://127.0.0.1:" + std::to_string(port) + "/mcp"; + metadata.authorization_servers = {"http://127.0.0.1:9000"}; + metadata.scopes_supported = {"mcp:read", "mcp:write"}; + manager.set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse document; + RawResponse denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + document = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, + "/.well-known/oauth-protected-resource/mcp"); + denied = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/mcp"); + manager.close(); + }, + asio::detached); + io_ctx_.run(); + + ASSERT_EQ(document.status, 200); + EXPECT_EQ(document.content_type, "application/json"); + const auto parsed = json::parse(document.body); + EXPECT_EQ(parsed.at("resource"), "http://127.0.0.1:" + std::to_string(port) + "/mcp"); + EXPECT_EQ(parsed.at("authorization_servers"), json::array({"http://127.0.0.1:9000"})); + EXPECT_EQ(parsed.at("scopes_supported"), json::array({"mcp:read", "mcp:write"})); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(denied.www_authenticate, R"(Bearer resource_metadata="http://127.0.0.1:)" + + std::to_string(port) + + R"(/.well-known/oauth-protected-resource/mcp")"); +} + +TEST_F(SessionManagerTest, MetadataRouteIsServedAheadOfTheCustomRequestHandler) { + const unsigned short port = 19133; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + manager.set_custom_request_handler( + [](const mcp::StringRequest& request) -> std::optional { + mcp::StringResponse response{http::status::ok, request.version()}; + response.body() = "from custom handler"; + response.prepare_payload(); + return response; + }); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "http://127.0.0.1:" + std::to_string(port) + "/mcp"; + manager.set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse document; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + document = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, + "/.well-known/oauth-protected-resource/mcp"); + manager.close(); + }, + asio::detached); + io_ctx_.run(); + + ASSERT_EQ(document.status, 200); + EXPECT_EQ(json::parse(document.body).at("resource"), + "http://127.0.0.1:" + std::to_string(port) + "/mcp"); +} + +TEST_F(SessionManagerTest, UnauthenticatedPathsBypassTheBearerCheck) { + const unsigned short port = 19134; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + manager.set_unauthenticated_paths({"/health"}); + manager.set_custom_request_handler( + [](const mcp::StringRequest& request) -> std::optional { + if (mcp::http_request_path(request.target()) != "/health") { + return std::nullopt; + } + mcp::StringResponse response{http::status::ok, request.version()}; + response.body() = "healthy"; + response.prepare_payload(); + return response; + }); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse exempt; + RawResponse exempt_with_query; + RawResponse guarded; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + exempt = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/health"); + exempt_with_query = + co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/health?probe=1"); + guarded = co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/healthy"); + manager.close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(exempt.status, 200); + EXPECT_EQ(exempt.body, "healthy"); + EXPECT_EQ(exempt_with_query.status, 200); + EXPECT_EQ(exempt_with_query.body, "healthy"); + EXPECT_EQ(guarded.status, 401); +} + +TEST_F(SessionManagerTest, RequestBodyBeyondTheLimitIsRejected) { + const unsigned short port = 19135; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + unsigned int status = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + // Announcing the length is enough: the parser rejects the request before the body is + // sent, which is the point of the limit. + const std::string header = + "POST /mcp HTTP/1.1\r\nHost: 127.0.0.1\r\nContent-Type: application/json\r\n" + "Content-Length: " + + std::to_string(mcp::constants::g_default_max_request_body_bytes + 1) + "\r\n\r\n"; + + beast::tcp_stream stream(io_ctx_.get_executor()); + asio::ip::tcp::resolver resolver(io_ctx_.get_executor()); + auto endpoints = + co_await resolver.async_resolve("127.0.0.1", std::to_string(port), asio::use_awaitable); + co_await stream.async_connect(*endpoints.begin(), asio::use_awaitable); + co_await asio::async_write(stream, asio::buffer(header), asio::use_awaitable); + + beast::flat_buffer response_buffer; + http::response response; + co_await http::async_read(stream, response_buffer, response, asio::use_awaitable); + status = response.result_int(); + + beast::error_code shutdown_error; + stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + }, + [&manager](const std::exception_ptr&) { manager.close(); }); + io_ctx_.run(); + + EXPECT_EQ(status, 413); +} + +// The two setters are order-independent: each renders from the whole current configuration, so +// neither call can strand the other's contribution. Both orders are exercised against real +// listeners and the resulting headers compared to each other. +TEST_F(SessionManagerTest, TheTwoChallengeSettersAreOrderIndependent) { + const auto run = [this](unsigned short port, bool metadata_first) -> std::string { + asio::io_context io_ctx; + mcp::StreamableHttpSessionManager manager(io_ctx.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "https://mcp.example.com/mcp"; + mcp::BearerChallengeConfig challenge; + challenge.realm = "mcp"; + challenge.scope = "mcp:read"; + + if (metadata_first) { + manager.set_protected_resource_metadata(metadata); + manager.set_bearer_challenge(challenge); + } else { + manager.set_bearer_challenge(challenge); + manager.set_protected_resource_metadata(metadata); + } + + asio::co_spawn(io_ctx, manager.listen(), asio::detached); + RawResponse denied; + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + denied = co_await raw_request(io_ctx.get_executor(), port, http::verb::get, "/mcp"); + manager.close(); + }, + asio::detached); + io_ctx.run(); + return denied.www_authenticate; + }; + + const auto metadata_first = run(19136, true); + const auto challenge_first = run(19137, false); + + EXPECT_EQ(metadata_first, challenge_first); + EXPECT_EQ( + metadata_first, + R"(Bearer realm="mcp", scope="mcp:read", )" + R"(resource_metadata="https://mcp.example.com/.well-known/oauth-protected-resource/mcp")"); +} + +// Exempting the path MCP is served on must not hand out unauthenticated MCP. The session count is +// the assertion that matters: a 404 that still created a session would have dispatched. +TEST_F(SessionManagerTest, ExemptingTheMcpPathRefusesToServeMcpUnauthenticated) { + const unsigned short port = 19138; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + manager.set_unauthenticated_paths({"/mcp"}); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse initialized; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + initialized = co_await do_initialize(io_ctx_.get_executor(), port); + manager.close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(initialized.status, 404); + EXPECT_TRUE(initialized.session_id.empty()); + EXPECT_EQ(manager.session_count(), 0U) << "an exempt path created an unauthenticated session"; +} + +TEST_F(SessionManagerTest, AsyncBearerValidatorDecidesWithoutBlocking) { + const unsigned short port = 19139; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + std::atomic validator_calls{0}; + manager.set_async_bearer_token_validator([&validator_calls](std::string token) -> mcp::Task { + validator_calls.fetch_add(1, std::memory_order_relaxed); + asio::steady_timer timer(co_await asio::this_coro::executor); + timer.expires_after(std::chrono::milliseconds(1)); + co_await timer.async_wait(asio::use_awaitable); + co_return token == "valid-token"; + }); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse denied; + RawResponse accepted; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto init_body = + json{{"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"clientInfo", {{"name", "test"}, {"version", "1"}}}, + {"capabilities", json::object()}}}, + {"id", 1}} + .dump(); + denied = co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + init_body, {}, std::string(mcp::g_LATEST_PROTOCOL_VERSION), + std::nullopt, "wrong-token"); + accepted = co_await raw_request(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + init_body, {}, std::string(mcp::g_LATEST_PROTOCOL_VERSION), + std::nullopt, "valid-token"); + manager.close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(accepted.status, 200); + EXPECT_EQ(validator_calls.load(std::memory_order_relaxed), 2); +} + +TEST_F(SessionManagerTest, OnlyOneBearerValidatorMayBeInstalled) { + const unsigned short port = 19140; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view) { return true; }); + EXPECT_THROW(manager.set_async_bearer_token_validator( + [](std::string) -> mcp::Task { co_return true; }), + std::logic_error); + manager.close(); +} + +TEST_F(SessionManagerTest, RequestBodyLimitIsConfigurable) { + const unsigned short port = 19141; + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_max_request_body_bytes(4096); + EXPECT_THROW(manager.set_max_request_body_bytes(0), std::invalid_argument); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + unsigned int status = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const std::string header = + "POST /mcp HTTP/1.1\r\nHost: 127.0.0.1\r\nContent-Type: application/json\r\n" + "Content-Length: 8192\r\n\r\n"; + + beast::tcp_stream stream(io_ctx_.get_executor()); + asio::ip::tcp::resolver resolver(io_ctx_.get_executor()); + auto endpoints = + co_await resolver.async_resolve("127.0.0.1", std::to_string(port), asio::use_awaitable); + co_await stream.async_connect(*endpoints.begin(), asio::use_awaitable); + co_await asio::async_write(stream, asio::buffer(header), asio::use_awaitable); + + beast::flat_buffer response_buffer; + http::response response; + co_await http::async_read(stream, response_buffer, response, asio::use_awaitable); + status = response.result_int(); + + beast::error_code shutdown_error; + stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + }, + [&manager](const std::exception_ptr&) { manager.close(); }); + io_ctx_.run(); + + // 8 KB clears the default but not the 4 KB cap installed above, so the cap is what answered. + EXPECT_EQ(status, 413); +} + +// =========================================================================== +// End to end: the SDK's own client completes discovery against the SDK's own server +// =========================================================================== + +// Without a metadata location in the challenge the SDK's client cannot be driven by the SDK's server: +// discovery has nothing to start from. This exercises the whole path on real sockets -- an +// unauthenticated request draws a 401, the challenge is parsed by the SDK's own parser, the URL it +// names is fetched by the SDK's own discovery client, and the document that comes back is the one the +// server was configured with. Nothing here is a string comparison against a hand-written header. +TEST_F(SessionManagerTest, SdkClientDiscoversAuthorizationFromTheSdkServersOwnChallenge) { + const unsigned short port = 19142; + const auto origin = "http://127.0.0.1:" + std::to_string(port); + const auto resource = origin + "/mcp"; + + mcp::StreamableHttpSessionManager manager(io_ctx_.get_executor(), "127.0.0.1", port, + make_echo_server_factory()); + manager.set_bearer_token_validator([](std::string_view token) { return token == "valid-token"; }); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = resource; + metadata.authorization_servers = {"https://auth.example.com"}; + metadata.scopes_supported = {"mcp:read"}; + manager.set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx_, manager.listen(), asio::detached); + + RawResponse challenge_response; + std::optional discovered; + std::string discovery_failure; + std::string advertised_metadata_url; + + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + // 1. An unauthenticated request draws the challenge. + challenge_response = + co_await raw_request(io_ctx_.get_executor(), port, http::verb::get, "/mcp"); + + // 2. The SDK's own parser reads it, exactly as OAuthAuthorizationManager would. + const auto challenges = + mcp::auth::parse_www_authenticate(challenge_response.www_authenticate); + const auto bearer = mcp::auth::select_bearer_challenge(challenges); + if (bearer.has_value() && bearer->resource_metadata.has_value()) { + advertised_metadata_url = *bearer->resource_metadata; + + // 3. The SDK's own discovery client fetches the URL the challenge named. + mcp::auth::MetadataFetchPolicy policy; + policy.allowed_origins.push_back(origin); + policy.allow_plain_http_loopback = true; + + auto http_client = + std::make_shared(co_await asio::this_coro::executor); + http_client->set_metadata_policy(policy); + http_client->set_host_resolver([](const std::string&, const std::string&) { + return std::vector{"127.0.0.1"}; + }); + mcp::auth::OAuthDiscoveryClient discovery(http_client); + try { + discovered = co_await discovery.discover_protected_resource( + resource, advertised_metadata_url); + } catch (const std::exception& error) { + discovery_failure = error.what(); + } + } + manager.close(); + }, + asio::detached); + + io_ctx_.run(); + + ASSERT_EQ(challenge_response.status, 401U); + EXPECT_EQ(advertised_metadata_url, origin + "/.well-known/oauth-protected-resource/mcp") + << "challenge was: " << challenge_response.www_authenticate; + + ASSERT_TRUE(discovered.has_value()) << "discovery failed: " << discovery_failure; + EXPECT_EQ(discovered->resource, resource); + ASSERT_EQ(discovered->authorization_servers.size(), 1U); + EXPECT_EQ(discovered->authorization_servers.front(), "https://auth.example.com"); + ASSERT_TRUE(discovered->scopes_supported.has_value()); + EXPECT_EQ(discovered->scopes_supported->front(), "mcp:read"); +} diff --git a/test/transport/transport_http_test.cpp b/test/transport/transport_http_test.cpp index e42e9bc..53d7154 100644 --- a/test/transport/transport_http_test.cpp +++ b/test/transport/transport_http_test.cpp @@ -3,14 +3,35 @@ #include +#include "../support/resolve_gate.hpp" +#include "../support/stalling_server.hpp" + +#include +#include #include #include +#include #include +#include +#include +#include +#include +#include #include +#include #include #include +#include +#include +#include +#include +#include +#include #include +#include #include +#include +#include namespace { @@ -33,6 +54,19 @@ TEST_F(HttpTransportTest, PortReturnsEphemeralPortWhenBoundToZero) { server_transport.close(); } +TEST_F(HttpTransportTest, PendingServerReadOwnsImplementationAfterTransportDestruction) { + auto server = std::make_unique(io_ctx_.get_executor(), "127.0.0.1", 0); + auto read = server->read_message(); + server.reset(); + + std::exception_ptr read_error; + asio::co_spawn(io_ctx_, std::move(read), + [&read_error](std::exception_ptr error, std::string) { read_error = error; }); + io_ctx_.run(); + + EXPECT_NE(read_error, nullptr); +} + TEST_F(HttpTransportTest, SendMessageServerReceives) { mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18080); @@ -230,57 +264,196 @@ TEST_F(HttpTransportTest, SessionIdPropagation) { EXPECT_TRUE(second_request_received); } -TEST_F(HttpTransportTest, ClientCloseDeletesSentToServer) { - mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18084); +namespace { - asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); +/// A raw HTTP peer on its own thread that hands out one MCP session and counts what a client +/// transport does to it. HttpServerTransport cannot stand in: it does not report a session DELETE. +/// A request made after initialize is answered only once the test releases it. +class SessionPeer final { + public: + SessionPeer() : acceptor_(io_ctx_, {asio::ip::make_address("127.0.0.1"), 0}), hold_(io_ctx_) { + hold_.expires_at(std::chrono::steady_clock::time_point::max()); + asio::co_spawn(io_ctx_, accept_loop(), asio::detached); + thread_ = std::thread([this]() { io_ctx_.run(); }); + } + + ~SessionPeer() { + io_ctx_.stop(); + thread_.join(); + } + + SessionPeer(const SessionPeer&) = delete; + SessionPeer& operator=(const SessionPeer&) = delete; + + [[nodiscard]] unsigned short port() const { return acceptor_.local_endpoint().port(); } + + [[nodiscard]] int accepted() { + std::lock_guard lock(mutex_); + return accepted_; + } + + [[nodiscard]] int deletes() { + std::lock_guard lock(mutex_); + return deletes_; + } + + [[nodiscard]] bool wait_for_deletes(int count, std::chrono::seconds limit) { + std::unique_lock lock(mutex_); + return changed_.wait_for(lock, limit, [this, count]() { return deletes_ >= count; }); + } + + [[nodiscard]] bool wait_for_held_request(std::chrono::seconds limit) { + std::unique_lock lock(mutex_); + return changed_.wait_for(lock, limit, [this]() { return held_requests_ >= 1; }); + } + + void release_held_response() { + asio::post(io_ctx_, [this]() { + released_ = true; + hold_.cancel(); + }); + } + + /// True once the released response has been written out in full. + [[nodiscard]] bool wait_for_held_response_sent(std::chrono::seconds limit) { + std::unique_lock lock(mutex_); + return changed_.wait_for(lock, limit, [this]() { return held_responses_sent_ >= 1; }); + } + + private: + void count(int& counter) { + { + std::lock_guard lock(mutex_); + ++counter; + } + changed_.notify_all(); + } + + mcp::Task accept_loop() { + for (;;) { + auto socket = co_await acceptor_.async_accept(asio::use_awaitable); + count(accepted_); + asio::co_spawn(io_ctx_, serve(std::move(socket)), asio::detached); + } + } + + mcp::Task serve(asio::ip::tcp::socket socket) { + try { + beast::flat_buffer buffer; + for (;;) { + http::request request; + co_await http::async_read(socket, buffer, request, asio::use_awaitable); + + http::response response{http::status::ok, request.version()}; + response.keep_alive(true); + bool held = false; + if (request.method() == http::verb::delete_) { + count(deletes_); + } else { + const auto message = nlohmann::json::parse(request.body()); + nlohmann::json reply = {{"jsonrpc", "2.0"}, + {"id", message.at("id")}, + {"result", nlohmann::json::object()}}; + const bool initialize = message.value("method", "") == "initialize"; + held = !initialize; + if (held) { + count(held_requests_); + while (!released_) { + boost::system::error_code ignored; + co_await hold_.async_wait( + asio::redirect_error(asio::use_awaitable, ignored)); + } + } + if (initialize) { + response.set("MCP-Session-Id", "session-1"); + reply["result"] = { + {"protocolVersion", mcp::g_LATEST_PROTOCOL_VERSION}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}; + } + response.set(http::field::content_type, "application/json"); + response.body() = reply.dump(); + } + response.prepare_payload(); + co_await http::async_write(socket, response, asio::use_awaitable); + if (held) { + count(held_responses_sent_); + } + } + } catch (const std::exception&) { + // The client went away; this connection is done. + (void)0; + } + } - bool session_terminated = false; - asio::co_spawn( - io_ctx_, - [&]() -> mcp::Task { - auto initialize_request = co_await server_transport.read_message(); - auto init_json = nlohmann::json::parse(initialize_request); + asio::io_context io_ctx_; + asio::ip::tcp::acceptor acceptor_; + /// Touched only on the peer's own thread. + asio::steady_timer hold_; + bool released_{false}; + std::thread thread_; + std::mutex mutex_; + std::condition_variable changed_; + int accepted_{0}; + int deletes_{0}; + int held_requests_{0}; + int held_responses_sent_{0}; +}; - nlohmann::json initialize_response = { - {"jsonrpc", "2.0"}, - {"result", - {{"protocolVersion", mcp::g_LATEST_PROTOCOL_VERSION}, - {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, - {"capabilities", {}}}}, - {"id", init_json.at("id")}}; +std::string initialize_request_text() { + const nlohmann::json initialize_request = { + {"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", mcp::g_LATEST_PROTOCOL_VERSION}, + {"clientInfo", {{"name", "test-client"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}}, + {"id", 1}}; + return initialize_request.dump(); +} - co_await server_transport.write_message(initialize_response.dump()); - session_terminated = true; - }, - asio::detached); +} // namespace + +TEST_F(HttpTransportTest, ClientCloseDeletesSentToServer) { + SessionPeer peer; + auto transport = std::make_shared( + io_ctx_.get_executor(), "http://127.0.0.1:" + std::to_string(peer.port()) + "/mcp"); + std::exception_ptr failure; + std::promise initialize_done; + auto initialized = initialize_done.get_future(); asio::co_spawn( io_ctx_, [&]() -> mcp::Task { - mcp::HttpClientTransport client_transport(io_ctx_.get_executor(), - "http://127.0.0.1:18084/mcp"); - - nlohmann::json initialize_request = { - {"jsonrpc", "2.0"}, - {"method", "initialize"}, - {"params", - {{"protocolVersion", mcp::g_LATEST_PROTOCOL_VERSION}, - {"clientInfo", {{"name", "test-client"}, {"version", "1.0.0"}}}, - {"capabilities", {}}}}, - {"id", 1}}; - - co_await client_transport.write_message(initialize_request.dump()); - co_await client_transport.read_message(); - - client_transport.close(); - server_transport.close(); + try { + co_await transport->write_message(initialize_request_text()); + co_await transport->read_message(); + } catch (...) { + failure = std::current_exception(); + } + initialize_done.set_value(); }, asio::detached); - io_ctx_.run(); - - EXPECT_TRUE(session_terminated); + auto work = asio::make_work_guard(io_ctx_); + std::thread runner([this]() { io_ctx_.run(); }); + + const auto limit = std::chrono::seconds(10); + const bool session_started = initialized.wait_for(limit) == std::future_status::ready; + bool delete_received = false; + if (session_started) { + transport->close(); + delete_received = peer.wait_for_deletes(1, limit); + } + + // Everything below reads state the runner thread wrote, so it stops first. + io_ctx_.stop(); + runner.join(); + + ASSERT_TRUE(session_started); + ASSERT_EQ(failure, nullptr); + EXPECT_TRUE(delete_received) << "close() did not send the session DELETE"; + EXPECT_EQ(peer.deletes(), 1); } TEST_F(HttpTransportTest, ServerRejectsInvalidProtocolVersion) { @@ -462,3 +635,2381 @@ TEST_F(HttpTransportTest, SessionRequestsAllowMissingNegotiatedProtocolHeader) { EXPECT_TRUE(tool_request_received); } + +TEST_F(HttpTransportTest, ClientUsesNegotiatedProtocolVersionForSessionRequests) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18088); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + bool second_request_received = false; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto initialize_request = co_await server_transport.read_message(); + const auto initialize_json = nlohmann::json::parse(initialize_request); + + const nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", "2025-06-18"}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", {}}}}, + {"id", initialize_json.at("id")}}; + co_await server_transport.write_message(initialize_response.dump()); + + const auto second_request = co_await server_transport.read_message(); + second_request_received = true; + const nlohmann::json second_response = { + {"jsonrpc", "2.0"}, + {"result", nlohmann::json::object()}, + {"id", nlohmann::json::parse(second_request).at("id")}}; + co_await server_transport.write_message(second_response.dump()); + }, + asio::detached); + + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + mcp::HttpClientTransport client_transport(io_ctx_.get_executor(), + "http://127.0.0.1:18088/mcp"); + const nlohmann::json initialize_request = { + {"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", "2025-06-18"}, + {"clientInfo", {{"name", "test-client"}, {"version", "1.0.0"}}}, + {"capabilities", {}}}}, + {"id", 1}}; + + co_await client_transport.write_message(initialize_request.dump()); + co_await client_transport.read_message(); + + const nlohmann::json subsequent_request = { + {"jsonrpc", "2.0"}, {"method", "tools/list"}, {"id", 2}}; + co_await client_transport.write_message(subsequent_request.dump()); + co_await client_transport.read_message(); + + client_transport.close(); + }, + [&server_transport](std::exception_ptr error) { + EXPECT_EQ(error, nullptr); + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_TRUE(second_request_received); +} + +TEST_F(HttpTransportTest, ClientCloseCancelsInFlightPostAfterSessionInitialization) { + struct TestState { + bool initialized{false}; + bool post_received{false}; + bool write_completed{false}; + bool write_cancelled{false}; + bool cleanup_completed{false}; + std::string write_outcome{"not completed"}; + std::exception_ptr server_error; + std::exception_ptr client_error; + }; + + auto server_transport = + std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 18089); + auto client_transport = std::make_shared(io_ctx_.get_executor(), + "http://127.0.0.1:18089/mcp"); + auto state = std::make_shared(); + auto post_received = std::make_shared(io_ctx_.get_executor()); + auto write_completed = std::make_shared(io_ctx_.get_executor()); + post_received->expires_at(std::chrono::steady_clock::time_point::max()); + write_completed->expires_at(std::chrono::steady_clock::time_point::max()); + + asio::co_spawn(io_ctx_, server_transport->listen(), asio::detached); + asio::co_spawn( + io_ctx_, + [](std::shared_ptr server, std::shared_ptr test_state, + std::shared_ptr post_signal) -> mcp::Task { + const auto initialize_message = co_await server->read_message(); + const auto initialize_request = nlohmann::json::parse(initialize_message); + const nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", mcp::g_LATEST_PROTOCOL_VERSION}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", {}}}}, + {"id", initialize_request.at("id")}}; + co_await server->write_message(initialize_response.dump()); + + const auto initialized_message = co_await server->read_message(); + EXPECT_EQ(nlohmann::json::parse(initialized_message).at("method"), + "notifications/initialized"); + + const auto pending_message = co_await server->read_message(); + EXPECT_EQ(nlohmann::json::parse(pending_message).at("method"), "tools/list"); + test_state->post_received = true; + post_signal->cancel(); + }(server_transport, state, post_received), + [state, post_received, server_transport](std::exception_ptr error) { + state->server_error = error; + if (error) { + post_received->cancel(); + server_transport->close(); + } + }); + + asio::co_spawn( + io_ctx_, + [](std::shared_ptr client, + std::shared_ptr server, std::shared_ptr test_state, + std::shared_ptr post_signal, + std::shared_ptr write_signal) -> mcp::Task { + const nlohmann::json initialize_request = { + {"jsonrpc", "2.0"}, + {"method", "initialize"}, + {"params", + {{"protocolVersion", mcp::g_LATEST_PROTOCOL_VERSION}, + {"clientInfo", {{"name", "test-client"}, {"version", "1.0.0"}}}, + {"capabilities", {}}}}, + {"id", 1}}; + co_await client->write_message(initialize_request.dump()); + const auto initialize_response = nlohmann::json::parse(co_await client->read_message()); + EXPECT_TRUE(initialize_response.contains("result")); + EXPECT_FALSE(client->session_id().empty()); + test_state->initialized = true; + + const nlohmann::json initialized_notification = {{"jsonrpc", "2.0"}, + {"method", "notifications/initialized"}}; + co_await client->write_message(initialized_notification.dump()); + + const nlohmann::json pending_request = { + {"jsonrpc", "2.0"}, {"method", "tools/list"}, {"id", 2}}; + asio::co_spawn(write_signal->get_executor(), client->write_message(pending_request.dump()), + [client, test_state, write_signal](std::exception_ptr error) { + test_state->write_completed = true; + test_state->write_outcome = "no error"; + if (error) { + try { + std::rethrow_exception(error); + } catch (const boost::system::system_error& system_error) { + test_state->write_cancelled = + system_error.code() == asio::error::operation_aborted; + test_state->write_outcome = + "system_error: " + system_error.code().message(); + } catch (const std::exception& other) { + test_state->write_outcome = other.what(); + } catch (...) { + test_state->write_outcome = "unknown exception"; + } + } + write_signal->cancel(); + }); + + if (!test_state->post_received) { + try { + co_await post_signal->async_wait(asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != asio::error::operation_aborted) { + throw; + } + } + } + if (test_state->server_error) { + std::rethrow_exception(test_state->server_error); + } + + client->close(); + if (!test_state->write_completed) { + try { + co_await write_signal->async_wait(asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != asio::error::operation_aborted) { + throw; + } + } + } + + asio::steady_timer cleanup_poll(write_signal->get_executor()); + while (!client->session_id().empty()) { + cleanup_poll.expires_after(std::chrono::milliseconds(1)); + co_await cleanup_poll.async_wait(asio::use_awaitable); + } + + server->close(); + test_state->cleanup_completed = true; + }(client_transport, server_transport, state, post_received, write_completed), + [state, client_transport, server_transport](std::exception_ptr error) { + state->client_error = error; + if (error) { + client_transport->close(); + server_transport->close(); + } + }); + + // Bound the test even if cancellation or listener cleanup regresses. + io_ctx_.run_for(std::chrono::seconds(5)); + const bool event_loop_drained = io_ctx_.stopped(); + if (!event_loop_drained) { + client_transport->close(); + server_transport->close(); + io_ctx_.stop(); + } + + EXPECT_EQ(state->server_error, nullptr); + EXPECT_EQ(state->client_error, nullptr); + EXPECT_TRUE(state->initialized); + EXPECT_TRUE(state->post_received); + EXPECT_TRUE(state->write_completed); + EXPECT_TRUE(state->write_cancelled) << "the write ended with: " << state->write_outcome; + EXPECT_TRUE(state->cleanup_completed); + EXPECT_TRUE(event_loop_drained); +} + +TEST_F(HttpTransportTest, PendingOperationsOwnImplementationAfterTransportDestruction) { + auto transport = std::make_unique(io_ctx_.get_executor(), + "http://127.0.0.1:18111/mcp"); + const nlohmann::json notification = {{"jsonrpc", "2.0"}, {"method", "notifications/initialized"}}; + + auto write = transport->write_message(notification.dump()); + auto read = transport->read_message(); + transport.reset(); + + std::exception_ptr write_error; + std::exception_ptr read_error; + asio::co_spawn(io_ctx_, std::move(write), + [&write_error](std::exception_ptr error) { write_error = error; }); + asio::co_spawn(io_ctx_, std::move(read), + [&read_error](std::exception_ptr error, std::string) { read_error = error; }); + + io_ctx_.run(); + + EXPECT_NE(write_error, nullptr); + EXPECT_NE(read_error, nullptr); +} + +TEST_F(HttpTransportTest, SerializesConcurrentWritesAcrossTwoIoThreads) { + constexpr int message_count = 64; + constexpr auto timeout = std::chrono::seconds(30); + + asio::io_context server_io; + auto server = + std::make_shared(server_io.get_executor(), "127.0.0.1", 18110); + auto server_deadline = std::make_shared(server_io.get_executor()); + server_deadline->expires_after(timeout); + auto received = std::make_shared>(message_count, false); + std::exception_ptr server_error; + + asio::co_spawn(server_io, server->listen(), asio::detached); + asio::co_spawn( + server_io, + [server, server_deadline, received]() -> mcp::Task { + for (int index = 0; index < message_count; ++index) { + const auto message = nlohmann::json::parse(co_await server->read_message()); + if (message.at("method") != "notifications/test" || !message.contains("params") || + !message.at("params").contains("sequence")) { + throw std::runtime_error("Unexpected concurrent-write test message"); + } + const auto sequence = message.at("params").at("sequence").get(); + if (sequence < 0 || sequence >= message_count || (*received)[sequence]) { + throw std::runtime_error("Invalid or duplicate concurrent-write sequence"); + } + (*received)[sequence] = true; + } + server_deadline->cancel(); + server->close(); + }, + [&server_error, server](std::exception_ptr error) { + server_error = error; + server->close(); + }); + server_deadline->async_wait([server](const boost::system::error_code& error) { + if (!error) { + server->close(); + } + }); + std::thread server_thread([&server_io] { server_io.run(); }); + + auto client = std::make_shared(io_ctx_.get_executor(), + "http://127.0.0.1:18110/mcp"); + auto coordination_strand = asio::make_strand(io_ctx_); + auto remaining = std::make_shared>(message_count); + auto completion_signal = std::make_shared(coordination_strand); + completion_signal->expires_at(std::chrono::steady_clock::time_point::max()); + auto client_deadline = std::make_shared(coordination_strand); + client_deadline->expires_after(timeout); + auto errors = std::make_shared>(); + auto errors_mutex = std::make_shared(); + + for (int index = 0; index < message_count; ++index) { + const nlohmann::json notification = { + {"jsonrpc", "2.0"}, {"method", "notifications/test"}, {"params", {{"sequence", index}}}}; + asio::co_spawn(io_ctx_, client->write_message(notification.dump()), + [remaining, completion_signal, errors, errors_mutex](std::exception_ptr error) { + if (error) { + std::lock_guard lock(*errors_mutex); + errors->push_back(error); + } + if (remaining->fetch_sub(1, std::memory_order_acq_rel) == 1) { + asio::post(completion_signal->get_executor(), + [completion_signal] { completion_signal->cancel(); }); + } + }); + } + + asio::co_spawn( + coordination_strand, + [client, completion_signal, client_deadline]() -> mcp::Task { + try { + co_await completion_signal->async_wait(asio::use_awaitable); + } catch (const boost::system::system_error& error) { + if (error.code() != asio::error::operation_aborted) { + throw; + } + } + client_deadline->cancel(); + client->close(); + }, + [client](std::exception_ptr) { client->close(); }); + client_deadline->async_wait([client, completion_signal](const boost::system::error_code& error) { + if (!error) { + client->close(); + completion_signal->cancel(); + } + }); + + std::thread first_client_thread([this] { io_ctx_.run(); }); + std::thread second_client_thread([this] { io_ctx_.run(); }); + first_client_thread.join(); + second_client_thread.join(); + server_thread.join(); + + EXPECT_EQ(server_error, nullptr); + EXPECT_TRUE(errors->empty()); + EXPECT_EQ(remaining->load(std::memory_order_acquire), 0); + EXPECT_TRUE(std::all_of(received->begin(), received->end(), [](bool value) { return value; })); +} + +TEST_F(HttpTransportTest, CloseCancelsIdleKeepAliveAndRejectsFurtherRequests) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + auto client = std::make_shared(io_ctx_.get_executor()); + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(2)); + + std::atomic authorization_calls{0}; + server->set_bearer_token_validator([&authorization_calls](std::string_view token) { + authorization_calls.fetch_add(1, std::memory_order_relaxed); + return token == "valid-token"; + }); + + bool first_request_completed = false; + bool connection_cancelled = false; + bool post_close_response_received = false; + bool timed_out = false; + std::exception_ptr client_error; + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + asio::co_spawn( + io_ctx_, + [server, client, deadline, &first_request_completed, &connection_cancelled, + &post_close_response_received]() -> mcp::Task { + asio::ip::tcp::resolver resolver(client->get_executor()); + auto endpoints = co_await resolver.async_resolve( + "127.0.0.1", std::to_string(server->port()), asio::use_awaitable); + co_await client->async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request first_request{http::verb::post, "/mcp", 11}; + first_request.set(http::field::host, "127.0.0.1"); + first_request.set(http::field::authorization, "Bearer valid-token"); + first_request.set(http::field::content_type, "application/json"); + first_request.keep_alive(true); + first_request.body() = R"({"jsonrpc":"2.0","method":"notifications/initialized"})"; + first_request.prepare_payload(); + co_await http::async_write(*client, first_request, asio::use_awaitable); + + beast::flat_buffer response_buffer; + http::response first_response; + co_await http::async_read(*client, response_buffer, first_response, asio::use_awaitable); + first_request_completed = first_response.result() == http::status::accepted; + + server->close(); + + http::request second_request{http::verb::post, "/mcp", 11}; + second_request.set(http::field::host, "127.0.0.1"); + second_request.set(http::field::authorization, "Bearer valid-token"); + second_request.set(http::field::content_type, "application/json"); + second_request.keep_alive(true); + second_request.body() = R"({"jsonrpc":"2.0","method":"notifications/post-close"})"; + second_request.prepare_payload(); + + try { + co_await http::async_write(*client, second_request, asio::use_awaitable); + http::response second_response; + co_await http::async_read(*client, response_buffer, second_response, + asio::use_awaitable); + post_close_response_received = true; + } catch (const boost::system::system_error&) { + connection_cancelled = true; + } + deadline->cancel(); + }, + [&client_error, deadline, server](std::exception_ptr error) { + client_error = std::move(error); + deadline->cancel(); + server->close(); + }); + deadline->async_wait([client, server, &timed_out](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server->close(); + boost::system::error_code ignored; + (void)client->socket().close(ignored); + }); + + io_ctx_.run(); + + EXPECT_EQ(client_error, nullptr); + EXPECT_FALSE(timed_out); + EXPECT_TRUE(first_request_completed); + EXPECT_TRUE(connection_cancelled); + EXPECT_FALSE(post_close_response_received); + EXPECT_EQ(authorization_calls.load(std::memory_order_relaxed), 1); +} + +TEST_F(HttpTransportTest, SessionlessDiscoverSucceedsAfterSessionEstablished) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18200); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(10)); + bool timed_out = false; + + bool discover_request_received = false; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto initialize_request = co_await server_transport.read_message(); + nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}}, + {"id", nlohmann::json::parse(initialize_request).at("id")}}; + co_await server_transport.write_message(initialize_response.dump()); + + const auto discover_request = co_await server_transport.read_message(); + discover_request_received = true; + nlohmann::json discover_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"supportedVersions", + nlohmann::json::array({std::string(mcp::g_PROTOCOL_VERSION_2026_07_28)})}}}, + {"id", nlohmann::json::parse(discover_request).at("id")}}; + co_await server_transport.write_message(discover_response.dump()); + }, + [](std::exception_ptr) {}); + + auto discover_status = http::status::unknown; + std::string discover_body; + std::string established_session_id; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18200", asio::use_awaitable); + + beast::tcp_stream init_stream(io_ctx_.get_executor()); + co_await init_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request init_request{http::verb::post, "/mcp", 11}; + init_request.set(http::field::host, "127.0.0.1"); + init_request.set(http::field::content_type, "application/json"); + init_request.body() = + R"({"jsonrpc":"2.0","method":"initialize","id":1,"params":{"protocolVersion":")" + + std::string(mcp::g_LATEST_PROTOCOL_VERSION) + + R"(","clientInfo":{"name":"test-client","version":"1.0.0"},"capabilities":{}}})"; + init_request.prepare_payload(); + co_await http::async_write(init_stream, init_request, asio::use_awaitable); + + beast::flat_buffer init_buffer; + http::response init_response; + co_await http::async_read(init_stream, init_buffer, init_response, asio::use_awaitable); + established_session_id = std::string(init_response["MCP-Session-Id"]); + + beast::error_code init_shutdown_error; + init_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, init_shutdown_error); + + // Fresh connection, no MCP-Session-Id header: server/discover must stay reachable + // even though the transport now holds an established session. + beast::tcp_stream discover_stream(io_ctx_.get_executor()); + co_await discover_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request discover_request{http::verb::post, "/mcp", 11}; + discover_request.set(http::field::host, "127.0.0.1"); + discover_request.set(http::field::content_type, "application/json"); + discover_request.body() = R"({"jsonrpc":"2.0","method":"server/discover","id":2})"; + discover_request.prepare_payload(); + co_await http::async_write(discover_stream, discover_request, asio::use_awaitable); + + beast::flat_buffer discover_buffer; + http::response discover_response; + co_await http::async_read(discover_stream, discover_buffer, discover_response, + asio::use_awaitable); + discover_status = discover_response.result(); + discover_body = discover_response.body(); + + beast::error_code discover_shutdown_error; + discover_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, + discover_shutdown_error); + + deadline->cancel(); + server_transport.close(); + }, + [](std::exception_ptr) {}); + + deadline->async_wait([&timed_out, &server_transport](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_FALSE(timed_out); + EXPECT_FALSE(established_session_id.empty()); + EXPECT_TRUE(discover_request_received); + EXPECT_EQ(discover_status, http::status::ok) << "body: " << discover_body; + // The caller must get its own id back. A sessionless request is carried internally under a + // transport-private id so it cannot squat the session's id space; that is an implementation + // detail the peer must never see. + EXPECT_EQ(nlohmann::json::parse(discover_body).at("id"), 2) << "body: " << discover_body; + // State the negative directly rather than inferring it from the id above: no part of the + // response a peer can read may carry the internal id, whatever shape it takes. + EXPECT_EQ(discover_body.find("mcp-pregate"), std::string::npos) << "body: " << discover_body; +} + +TEST_F(HttpTransportTest, SessionlessNonDiscoverRequestRejectedAfterSessionEstablished) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18201); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(10)); + bool timed_out = false; + + bool tools_list_reached_server = false; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto initialize_request = co_await server_transport.read_message(); + nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}}, + {"id", nlohmann::json::parse(initialize_request).at("id")}}; + co_await server_transport.write_message(initialize_response.dump()); + + co_await server_transport.read_message(); + tools_list_reached_server = true; + }, + [](std::exception_ptr) {}); + + auto tools_list_status = http::status::unknown; + std::string tools_list_body; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18201", asio::use_awaitable); + + beast::tcp_stream init_stream(io_ctx_.get_executor()); + co_await init_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request init_request{http::verb::post, "/mcp", 11}; + init_request.set(http::field::host, "127.0.0.1"); + init_request.set(http::field::content_type, "application/json"); + init_request.body() = + R"({"jsonrpc":"2.0","method":"initialize","id":1,"params":{"protocolVersion":")" + + std::string(mcp::g_LATEST_PROTOCOL_VERSION) + + R"(","clientInfo":{"name":"test-client","version":"1.0.0"},"capabilities":{}}})"; + init_request.prepare_payload(); + co_await http::async_write(init_stream, init_request, asio::use_awaitable); + + beast::flat_buffer init_buffer; + http::response init_response; + co_await http::async_read(init_stream, init_buffer, init_response, asio::use_awaitable); + + beast::error_code init_shutdown_error; + init_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, init_shutdown_error); + + // Same shape as the discover probe above, but a non-pre-gate method: the session + // gate must still reject it, pinning the discover exemption as method-specific. + beast::tcp_stream tools_stream(io_ctx_.get_executor()); + co_await tools_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request tools_request{http::verb::post, "/mcp", 11}; + tools_request.set(http::field::host, "127.0.0.1"); + tools_request.set(http::field::content_type, "application/json"); + tools_request.body() = R"({"jsonrpc":"2.0","method":"tools/list","id":2})"; + tools_request.prepare_payload(); + co_await http::async_write(tools_stream, tools_request, asio::use_awaitable); + + beast::flat_buffer tools_buffer; + http::response tools_response; + co_await http::async_read(tools_stream, tools_buffer, tools_response, asio::use_awaitable); + tools_list_status = tools_response.result(); + tools_list_body = tools_response.body(); + + beast::error_code tools_shutdown_error; + tools_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, tools_shutdown_error); + + deadline->cancel(); + server_transport.close(); + }, + [](std::exception_ptr) {}); + + deadline->async_wait([&timed_out, &server_transport](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_FALSE(timed_out); + EXPECT_FALSE(tools_list_reached_server); + EXPECT_EQ(tools_list_status, http::status::bad_request) << "body: " << tools_list_body; + EXPECT_NE(tools_list_body.find("Session active"), std::string::npos) << "body: " << tools_list_body; +} + +TEST_F(HttpTransportTest, DiscoverWithWrongSessionHeaderIsStillRejected) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18203); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(10)); + bool timed_out = false; + + bool discover_reached_server = false; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto initialize_request = co_await server_transport.read_message(); + nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}}, + {"id", nlohmann::json::parse(initialize_request).at("id")}}; + co_await server_transport.write_message(initialize_response.dump()); + + co_await server_transport.read_message(); + discover_reached_server = true; + }, + [](std::exception_ptr) {}); + + auto discover_status = http::status::unknown; + std::string discover_body; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18203", asio::use_awaitable); + + beast::tcp_stream init_stream(io_ctx_.get_executor()); + co_await init_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request init_request{http::verb::post, "/mcp", 11}; + init_request.set(http::field::host, "127.0.0.1"); + init_request.set(http::field::content_type, "application/json"); + init_request.body() = + R"({"jsonrpc":"2.0","method":"initialize","id":1,"params":{"protocolVersion":")" + + std::string(mcp::g_LATEST_PROTOCOL_VERSION) + + R"(","clientInfo":{"name":"test-client","version":"1.0.0"},"capabilities":{}}})"; + init_request.prepare_payload(); + co_await http::async_write(init_stream, init_request, asio::use_awaitable); + + beast::flat_buffer init_buffer; + http::response init_response; + co_await http::async_read(init_stream, init_buffer, init_response, asio::use_awaitable); + const auto established_session_id = std::string(init_response["MCP-Session-Id"]); + EXPECT_FALSE(established_session_id.empty()); + + beast::error_code init_shutdown_error; + init_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, init_shutdown_error); + + // The exemption is conditional on the session header being ABSENT. A discover + // request that presents a header, and presents the wrong one, must keep going + // through validate_post_session and be rejected exactly as before. + beast::tcp_stream discover_stream(io_ctx_.get_executor()); + co_await discover_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request discover_request{http::verb::post, "/mcp", 11}; + discover_request.set(http::field::host, "127.0.0.1"); + discover_request.set(http::field::content_type, "application/json"); + discover_request.set("MCP-Session-Id", established_session_id + "-tampered"); + discover_request.body() = R"({"jsonrpc":"2.0","method":"server/discover","id":2})"; + discover_request.prepare_payload(); + co_await http::async_write(discover_stream, discover_request, asio::use_awaitable); + + beast::flat_buffer discover_buffer; + http::response discover_response; + co_await http::async_read(discover_stream, discover_buffer, discover_response, + asio::use_awaitable); + discover_status = discover_response.result(); + discover_body = discover_response.body(); + + beast::error_code discover_shutdown_error; + discover_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, + discover_shutdown_error); + + deadline->cancel(); + server_transport.close(); + }, + [](std::exception_ptr) {}); + + deadline->async_wait([&timed_out, &server_transport](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_FALSE(timed_out); + EXPECT_FALSE(discover_reached_server); + EXPECT_EQ(discover_status, http::status::bad_request) << "body: " << discover_body; + EXPECT_NE(discover_body.find("Invalid MCP-Session-Id header"), std::string::npos) + << "body: " << discover_body; +} + +TEST_F(HttpTransportTest, SessionlessDiscoverNotificationIsStillRejected) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18204); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(10)); + bool timed_out = false; + + bool notification_reached_server = false; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto initialize_request = co_await server_transport.read_message(); + nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}}, + {"id", nlohmann::json::parse(initialize_request).at("id")}}; + co_await server_transport.write_message(initialize_response.dump()); + + co_await server_transport.read_message(); + notification_reached_server = true; + }, + [](std::exception_ptr) {}); + + auto notification_status = http::status::unknown; + std::string notification_body; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18204", asio::use_awaitable); + + beast::tcp_stream init_stream(io_ctx_.get_executor()); + co_await init_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request init_request{http::verb::post, "/mcp", 11}; + init_request.set(http::field::host, "127.0.0.1"); + init_request.set(http::field::content_type, "application/json"); + init_request.body() = + R"({"jsonrpc":"2.0","method":"initialize","id":1,"params":{"protocolVersion":")" + + std::string(mcp::g_LATEST_PROTOCOL_VERSION) + + R"(","clientInfo":{"name":"test-client","version":"1.0.0"},"capabilities":{}}})"; + init_request.prepare_payload(); + co_await http::async_write(init_stream, init_request, asio::use_awaitable); + + beast::flat_buffer init_buffer; + http::response init_response; + co_await http::async_read(init_stream, init_buffer, init_response, asio::use_awaitable); + + beast::error_code init_shutdown_error; + init_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, init_shutdown_error); + + // A sessionless server/discover with no "id" is a notification, not a request. It + // must NOT take the exemption, because doing so would push its body onto the + // unbounded incoming queue and answer 202 without the session gate ever running. + beast::tcp_stream notification_stream(io_ctx_.get_executor()); + co_await notification_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request notification_request{http::verb::post, "/mcp", 11}; + notification_request.set(http::field::host, "127.0.0.1"); + notification_request.set(http::field::content_type, "application/json"); + notification_request.body() = R"({"jsonrpc":"2.0","method":"server/discover"})"; + notification_request.prepare_payload(); + co_await http::async_write(notification_stream, notification_request, asio::use_awaitable); + + beast::flat_buffer notification_buffer; + http::response notification_response; + co_await http::async_read(notification_stream, notification_buffer, notification_response, + asio::use_awaitable); + notification_status = notification_response.result(); + notification_body = notification_response.body(); + + beast::error_code notification_shutdown_error; + notification_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, + notification_shutdown_error); + + deadline->cancel(); + server_transport.close(); + }, + [](std::exception_ptr) {}); + + deadline->async_wait([&timed_out, &server_transport](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_FALSE(timed_out); + EXPECT_FALSE(notification_reached_server); + EXPECT_EQ(notification_status, http::status::bad_request) << "body: " << notification_body; + EXPECT_NE(notification_body.find("Session active"), std::string::npos) + << "body: " << notification_body; +} + +// Sessionless server/discover is unauthenticated by construction, so its responses must never +// be appended to the replay EventStore. That store is a bounded ring shared with the +// established session: a flood of sessionless requests would evict the session's replay +// history, and its next Last-Event-ID resume would fail with 410 Gone, losing messages +// unrecoverably. The store is sized to 2 here so that three sessionless requests WOULD evict +// event "1" if they were stored, mirroring HttpResumabilityTest.GetWithEvictedEventIdReturns410 +// which produces exactly that 410 using the session's own traffic. +TEST_F(HttpTransportTest, SessionlessDiscoverDoesNotEvictSessionReplayHistory) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18205, 2); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(10)); + bool timed_out = false; + + constexpr int k_sessionless_request_count = 3; + + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto initialize_request = co_await server_transport.read_message(); + nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}}, + {"id", nlohmann::json::parse(initialize_request).at("id")}}; + co_await server_transport.write_message(initialize_response.dump()); + + for (int i = 0; i < k_sessionless_request_count; ++i) { + const auto discover_request = co_await server_transport.read_message(); + nlohmann::json discover_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"supportedVersions", + nlohmann::json::array({std::string(mcp::g_PROTOCOL_VERSION_2026_07_28)})}}}, + {"id", nlohmann::json::parse(discover_request).at("id")}}; + co_await server_transport.write_message(discover_response.dump()); + } + }, + [](std::exception_ptr) {}); + + int resume_status = 0; + int sessionless_ok_count = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18205", asio::use_awaitable); + + beast::tcp_stream init_stream(io_ctx_.get_executor()); + co_await init_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request init_request{http::verb::post, "/mcp", 11}; + init_request.set(http::field::host, "127.0.0.1"); + init_request.set(http::field::content_type, "application/json"); + init_request.body() = + R"({"jsonrpc":"2.0","method":"initialize","id":1,"params":{"protocolVersion":")" + + std::string(mcp::g_LATEST_PROTOCOL_VERSION) + + R"(","clientInfo":{"name":"test-client","version":"1.0.0"},"capabilities":{}}})"; + init_request.prepare_payload(); + co_await http::async_write(init_stream, init_request, asio::use_awaitable); + + beast::flat_buffer init_buffer; + http::response init_response; + co_await http::async_read(init_stream, init_buffer, init_response, asio::use_awaitable); + const auto established_session_id = std::string(init_response["MCP-Session-Id"]); + + beast::error_code init_shutdown_error; + init_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, init_shutdown_error); + + // Unauthenticated traffic: none of these carry an MCP-Session-Id header. + for (int i = 0; i < k_sessionless_request_count; ++i) { + beast::tcp_stream discover_stream(io_ctx_.get_executor()); + co_await discover_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request discover_request{http::verb::post, "/mcp", 11}; + discover_request.set(http::field::host, "127.0.0.1"); + discover_request.set(http::field::content_type, "application/json"); + discover_request.body() = R"({"jsonrpc":"2.0","method":"server/discover","id":)" + + std::to_string(101 + i) + "}"; + discover_request.prepare_payload(); + co_await http::async_write(discover_stream, discover_request, asio::use_awaitable); + + beast::flat_buffer discover_buffer; + http::response discover_response; + co_await http::async_read(discover_stream, discover_buffer, discover_response, + asio::use_awaitable); + if (discover_response.result() == http::status::ok) { + ++sessionless_ok_count; + } + + beast::error_code discover_shutdown_error; + discover_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, + discover_shutdown_error); + } + + // The established session resumes from the event id it already holds. + beast::tcp_stream resume_stream(io_ctx_.get_executor()); + co_await resume_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request resume_request{http::verb::get, "/mcp", 11}; + resume_request.set(http::field::host, "127.0.0.1"); + resume_request.set(http::field::accept, "text/event-stream"); + resume_request.set("MCP-Protocol-Version", std::string(mcp::g_LATEST_PROTOCOL_VERSION)); + resume_request.set("MCP-Session-Id", established_session_id); + resume_request.set("Last-Event-ID", "1"); + co_await http::async_write(resume_stream, resume_request, asio::use_awaitable); + + beast::flat_buffer resume_buffer; + http::response resume_response; + co_await http::async_read(resume_stream, resume_buffer, resume_response, + asio::use_awaitable); + resume_status = static_cast(resume_response.result_int()); + + beast::error_code resume_shutdown_error; + resume_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, + resume_shutdown_error); + + deadline->cancel(); + server_transport.close(); + }, + [](std::exception_ptr) {}); + + deadline->async_wait([&timed_out, &server_transport](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_FALSE(timed_out); + // All three were served, so a pass here is not an artifact of rejecting them. + EXPECT_EQ(sessionless_ok_count, k_sessionless_request_count); + EXPECT_EQ(resume_status, 200); + // Only the initialize response belongs in the store. + EXPECT_EQ(server_transport.event_store().size(), 1u); +} + +TEST_F(HttpTransportTest, DiscoverAcceptsDiscoverableOnlyProtocolVersionHeader) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18202); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(10)); + bool timed_out = false; + + bool discover_request_received = false; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto discover_request = co_await server_transport.read_message(); + discover_request_received = true; + nlohmann::json discover_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"supportedVersions", + nlohmann::json::array({std::string(mcp::g_PROTOCOL_VERSION_2026_07_28)})}}}, + {"id", nlohmann::json::parse(discover_request).at("id")}}; + co_await server_transport.write_message(discover_response.dump()); + }, + [](std::exception_ptr) {}); + + auto discover_status = http::status::unknown; + std::string discover_body; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18202", asio::use_awaitable); + + beast::tcp_stream stream(io_ctx_.get_executor()); + co_await stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request discover_request{http::verb::post, "/mcp", 11}; + discover_request.set(http::field::host, "127.0.0.1"); + discover_request.set(http::field::content_type, "application/json"); + // 2026-07-28 is discoverable but not in g_SUPPORTED_PROTOCOL_VERSIONS, so the + // negotiated-version check would reject it for any non-pre-gate method. + discover_request.set("MCP-Protocol-Version", + std::string(mcp::g_PROTOCOL_VERSION_2026_07_28)); + discover_request.body() = R"({"jsonrpc":"2.0","method":"server/discover","id":1})"; + discover_request.prepare_payload(); + co_await http::async_write(stream, discover_request, asio::use_awaitable); + + beast::flat_buffer response_buffer; + http::response discover_response; + co_await http::async_read(stream, response_buffer, discover_response, asio::use_awaitable); + discover_status = discover_response.result(); + discover_body = discover_response.body(); + + beast::error_code shutdown_error; + stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + + deadline->cancel(); + server_transport.close(); + }, + [](std::exception_ptr) {}); + + deadline->async_wait([&timed_out, &server_transport](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_FALSE(timed_out); + EXPECT_TRUE(discover_request_received); + EXPECT_EQ(discover_status, http::status::ok) << "body: " << discover_body; +} + +TEST_F(HttpTransportTest, SessionlessDiscoverCannotSquatSessionRequestId) { + mcp::HttpServerTransport server_transport(io_ctx_.get_executor(), "127.0.0.1", 18206); + + asio::co_spawn(io_ctx_, server_transport.listen(), asio::detached); + + auto deadline = std::make_shared(io_ctx_.get_executor()); + deadline->expires_after(std::chrono::seconds(10)); + bool timed_out = false; + + // Signalled once initialize has produced a session, so the sessionless prober runs while a + // session is live. + auto session_ready = std::make_shared(io_ctx_.get_executor()); + session_ready->expires_at(std::chrono::steady_clock::time_point::max()); + // Signalled once the sessionless discover has reached the server. Registration precedes the + // enqueue, so a discover that the server can read is a discover already holding its id. + auto discover_registered = std::make_shared(io_ctx_.get_executor()); + discover_registered->expires_at(std::chrono::steady_clock::time_point::max()); + + bool discover_reached_server = false; + bool session_request_reached_server = false; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + const auto initialize_request = co_await server_transport.read_message(); + nlohmann::json initialize_response = { + {"jsonrpc", "2.0"}, + {"result", + {{"protocolVersion", std::string(mcp::g_LATEST_PROTOCOL_VERSION)}, + {"serverInfo", {{"name", "test-server"}, {"version", "1.0.0"}}}, + {"capabilities", nlohmann::json::object()}}}, + {"id", nlohmann::json::parse(initialize_request).at("id")}}; + co_await server_transport.write_message(initialize_response.dump()); + + // The sessionless discover is deliberately left unanswered: it holds its pending + // entry for the whole test, which is what gives the session request something to + // collide with. + const auto discover_request = co_await server_transport.read_message(); + discover_reached_server = true; + discover_registered->cancel(); + + const auto session_request = co_await server_transport.read_message(); + session_request_reached_server = true; + nlohmann::json session_response = {{"jsonrpc", "2.0"}, + {"result", {{"tools", nlohmann::json::array()}}}, + {"id", nlohmann::json::parse(session_request).at("id")}}; + co_await server_transport.write_message(session_response.dump()); + }, + [](std::exception_ptr) {}); + + // The sessionless prober. It writes its discover and never reads the reply, so the pending + // entry it claims stays claimed. + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + try { + co_await session_ready->async_wait(asio::use_awaitable); + } catch (const boost::system::system_error&) { + // Cancelled: the session is up. + } + + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18206", asio::use_awaitable); + + auto probe_stream = std::make_shared(io_ctx_.get_executor()); + co_await probe_stream->async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request discover_request{http::verb::post, "/mcp", 11}; + discover_request.set(http::field::host, "127.0.0.1"); + discover_request.set(http::field::content_type, "application/json"); + discover_request.body() = R"({"jsonrpc":"2.0","method":"server/discover","id":7})"; + discover_request.prepare_payload(); + co_await http::async_write(*probe_stream, discover_request, asio::use_awaitable); + + // Hold the connection open for the rest of the test. + try { + co_await deadline->async_wait(asio::use_awaitable); + } catch (const boost::system::system_error&) { + // Cancelled at teardown. + } + }, + [](std::exception_ptr) {}); + + auto collision_status = http::status::unknown; + std::string collision_body; + std::string established_session_id; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto resolver = asio::ip::tcp::resolver(io_ctx_.get_executor()); + const auto endpoints = + co_await resolver.async_resolve("127.0.0.1", "18206", asio::use_awaitable); + + beast::tcp_stream session_stream(io_ctx_.get_executor()); + co_await session_stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request init_request{http::verb::post, "/mcp", 11}; + init_request.set(http::field::host, "127.0.0.1"); + init_request.set(http::field::content_type, "application/json"); + init_request.body() = + R"({"jsonrpc":"2.0","method":"initialize","id":1,"params":{"protocolVersion":")" + + std::string(mcp::g_LATEST_PROTOCOL_VERSION) + + R"(","clientInfo":{"name":"test-client","version":"1.0.0"},"capabilities":{}}})"; + init_request.prepare_payload(); + co_await http::async_write(session_stream, init_request, asio::use_awaitable); + + beast::flat_buffer init_buffer; + http::response init_response; + co_await http::async_read(session_stream, init_buffer, init_response, asio::use_awaitable); + established_session_id = std::string(init_response["MCP-Session-Id"]); + + session_ready->cancel(); + try { + co_await discover_registered->async_wait(asio::use_awaitable); + } catch (const boost::system::system_error&) { + // Cancelled: the sessionless discover holds id 7. + } + + // A legitimate, session-authenticated request that happens to reuse id 7. The + // sessionless prober picked that id out of a space the session owns, so this must + // still be served. + http::request tools_request{http::verb::post, "/mcp", 11}; + tools_request.set(http::field::host, "127.0.0.1"); + tools_request.set(http::field::content_type, "application/json"); + tools_request.set("MCP-Session-Id", established_session_id); + tools_request.body() = R"({"jsonrpc":"2.0","method":"tools/list","id":7})"; + tools_request.prepare_payload(); + co_await http::async_write(session_stream, tools_request, asio::use_awaitable); + + beast::flat_buffer tools_buffer; + http::response tools_response; + co_await http::async_read(session_stream, tools_buffer, tools_response, + asio::use_awaitable); + collision_status = tools_response.result(); + collision_body = tools_response.body(); + + beast::error_code shutdown_error; + session_stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + + deadline->cancel(); + server_transport.close(); + }, + [](std::exception_ptr) {}); + + deadline->async_wait([&timed_out, &server_transport](const boost::system::error_code& error) { + if (error) { + return; + } + timed_out = true; + server_transport.close(); + }); + + io_ctx_.run(); + + EXPECT_FALSE(timed_out); + EXPECT_FALSE(established_session_id.empty()); + EXPECT_TRUE(discover_reached_server); + EXPECT_TRUE(session_request_reached_server) + << "the session request never reached the server: its id was squatted"; + EXPECT_EQ(collision_status, http::status::ok) << "body: " << collision_body; + EXPECT_NE(collision_body.find("\"result\""), std::string::npos) << "body: " << collision_body; + // Replay-store classification must survive the containment. sessionless_request_ids decides + // replay-store exclusion, so if the prober's chosen id were still the key, the session's OWN + // response -- which carries that same id 7 -- would be misclassified as sessionless and silently + // dropped from the replay store. Both the initialize response and the id-7 session response + // belong there; only the unanswered sessionless discover does not. + EXPECT_EQ(server_transport.event_store().size(), 2u); +} + +TEST_F(HttpTransportTest, NonAtomicConfigurationLocksWhenListeningStarts) { + mcp::HttpServerTransport server(io_ctx_.get_executor(), "127.0.0.1", 0); + auto listener = server.listen(); + + EXPECT_NO_THROW(server.set_json_only(true)); + EXPECT_THROW(server.set_allowed_origins({"https://trusted.example"}), std::logic_error); + EXPECT_THROW(server.set_allow_all_origins(true), std::logic_error); + EXPECT_THROW(server.set_bearer_token_validator({}), std::logic_error); + EXPECT_THROW(server.set_bearer_challenge({}), std::logic_error); + EXPECT_THROW(server.set_protected_resource_metadata({}), std::logic_error); + EXPECT_THROW(server.set_unauthenticated_paths({"/health"}), std::logic_error); + EXPECT_THROW(server.set_async_bearer_token_validator({}), std::logic_error); + EXPECT_THROW(server.set_max_request_body_bytes(4096), std::logic_error); + + server.close(); + asio::co_spawn(io_ctx_, std::move(listener), asio::detached); + io_ctx_.run(); +} + +// A bearer provider installed after write_message() returns belongs to the next request, not to the +// one that call already started. HttpClientTransport pins the provider when the request begins, the +// way OAuthHttpClient::make_exchange() pins its per-exchange state, so the property holds without +// any timing window to hit: the swap below happens strictly after write_message() has returned and +// strictly before the request reaches the wire. +TEST_F(HttpTransportTest, WriteMessagePinsBearerProviderAtRequestStart) { + asio::ip::tcp::acceptor acceptor(io_ctx_, + asio::ip::tcp::endpoint(asio::ip::make_address("127.0.0.1"), 0)); + const auto port = acceptor.local_endpoint().port(); + + std::string observed_authorization; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + beast::tcp_stream stream(co_await acceptor.async_accept(asio::use_awaitable)); + + beast::flat_buffer buffer; + http::request request; + co_await http::async_read(stream, buffer, request, asio::use_awaitable); + observed_authorization = std::string(request[http::field::authorization]); + + http::response response{http::status::accepted, 11}; + response.prepare_payload(); + co_await http::async_write(stream, response, asio::use_awaitable); + + beast::error_code ignored; + stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, ignored); + }, + asio::detached); + + mcp::HttpClientTransport client(io_ctx_.get_executor(), + "http://127.0.0.1:" + std::to_string(port) + "/mcp"); + client.set_bearer_token_provider([]() { return std::string("pinned-at-start"); }); + + nlohmann::json notification = {{"jsonrpc", "2.0"}, {"method", "notifications/initialized"}}; + auto write = client.write_message(notification.dump()); + + client.set_bearer_token_provider([]() { return std::string("swapped-after-start"); }); + + std::exception_ptr write_error; + asio::co_spawn(io_ctx_, std::move(write), + [&write_error](std::exception_ptr error) { write_error = error; }); + io_ctx_.run(); + + ASSERT_EQ(write_error, nullptr); + EXPECT_EQ(observed_authorization, "Bearer pinned-at-start"); +} + +// =========================================================================== +// WWW-Authenticate challenge rendering +// =========================================================================== + +TEST(BearerChallengeTest, EmptyConfigRendersBareBearer) { + EXPECT_EQ(mcp::format_www_authenticate(mcp::BearerChallengeConfig{}), "Bearer"); +} + +TEST(BearerChallengeTest, EachParameterRendersOnItsOwn) { + mcp::BearerChallengeConfig realm_only; + realm_only.realm = "mcp"; + EXPECT_EQ(mcp::format_www_authenticate(realm_only), R"(Bearer realm="mcp")"); + + mcp::BearerChallengeConfig error_only; + error_only.error = "invalid_token"; + EXPECT_EQ(mcp::format_www_authenticate(error_only), R"(Bearer error="invalid_token")"); + + mcp::BearerChallengeConfig scope_only; + scope_only.scope = "mcp:read mcp:write"; + EXPECT_EQ(mcp::format_www_authenticate(scope_only), R"(Bearer scope="mcp:read mcp:write")"); + + mcp::BearerChallengeConfig metadata_only; + metadata_only.resource_metadata = "http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp"; + EXPECT_EQ( + mcp::format_www_authenticate(metadata_only), + R"(Bearer resource_metadata="http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp")"); +} + +TEST(BearerChallengeTest, AllParametersRenderInDocumentedOrder) { + mcp::BearerChallengeConfig challenge; + challenge.resource_metadata = "http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp"; + challenge.scope = "mcp:read"; + challenge.realm = "mcp"; + challenge.error = "invalid_token"; + + EXPECT_EQ(mcp::format_www_authenticate(challenge), + R"(Bearer realm="mcp", error="invalid_token", scope="mcp:read", )" + R"(resource_metadata="http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp")"); +} + +TEST(BearerChallengeTest, OrderIsIndependentOfAssignmentOrder) { + mcp::BearerChallengeConfig assigned_forwards; + assigned_forwards.realm = "r"; + assigned_forwards.scope = "s"; + + mcp::BearerChallengeConfig assigned_backwards; + assigned_backwards.scope = "s"; + assigned_backwards.realm = "r"; + + EXPECT_EQ(mcp::format_www_authenticate(assigned_forwards), R"(Bearer realm="r", scope="s")"); + EXPECT_EQ(mcp::format_www_authenticate(assigned_backwards), + mcp::format_www_authenticate(assigned_forwards)); +} + +TEST(BearerChallengeTest, BackslashAndQuoteAreEscaped) { + mcp::BearerChallengeConfig challenge; + challenge.realm = R"(a"b\c)"; + + EXPECT_EQ(mcp::format_www_authenticate(challenge), R"(Bearer realm="a\"b\\c")"); +} + +TEST(BearerChallengeTest, UnquotableValueIsRejected) { + mcp::BearerChallengeConfig carriage_return; + carriage_return.realm = "mcp\r\nX-Injected: 1"; + EXPECT_THROW(mcp::format_www_authenticate(carriage_return), std::invalid_argument); + + mcp::BearerChallengeConfig non_ascii; + non_ascii.scope = + "mcp:r\xc3\xa9" + "ad"; + EXPECT_THROW(mcp::format_www_authenticate(non_ascii), std::invalid_argument); +} + +TEST(HttpRequestPathTest, QueryAndFragmentAreStripped) { + EXPECT_EQ(mcp::http_request_path("/health"), "/health"); + EXPECT_EQ(mcp::http_request_path("/health?probe=1"), "/health"); + EXPECT_EQ(mcp::http_request_path("/health#frag"), "/health"); + EXPECT_EQ(mcp::http_request_path("/health?a=1#frag"), "/health"); +} + +TEST(ProtectedResourceMetadataTest, PathInsertsTheWellKnownSegmentBeforeTheResourcePath) { + // RFC 9728 3.1: a resource with a path is described under that path, not at the bare + // well-known location. Publishing at the bare path would leave clients looking elsewhere. + mcp::ProtectedResourceMetadataConfig with_path; + with_path.resource = "https://h/mcp"; + EXPECT_EQ(mcp::protected_resource_metadata_path(with_path), + "/.well-known/oauth-protected-resource/mcp"); + EXPECT_EQ(mcp::protected_resource_metadata_url(with_path), + "https://h/.well-known/oauth-protected-resource/mcp"); + + mcp::ProtectedResourceMetadataConfig nested; + nested.resource = "https://h/a/b"; + EXPECT_EQ(mcp::protected_resource_metadata_path(nested), + "/.well-known/oauth-protected-resource/a/b"); +} + +TEST(ProtectedResourceMetadataTest, PathIsBareOnlyForAResourceAtTheOriginRoot) { + mcp::ProtectedResourceMetadataConfig root; + root.resource = "https://h"; + EXPECT_EQ(mcp::protected_resource_metadata_path(root), "/.well-known/oauth-protected-resource"); + + mcp::ProtectedResourceMetadataConfig trailing_slash; + trailing_slash.resource = "https://h/"; + EXPECT_EQ(mcp::protected_resource_metadata_path(trailing_slash), + "/.well-known/oauth-protected-resource"); +} + +TEST(ProtectedResourceMetadataTest, AnExplicitPathOverridesTheDerivation) { + mcp::ProtectedResourceMetadataConfig overridden; + overridden.resource = "https://h/mcp"; + overridden.path = "/custom-metadata"; + + EXPECT_EQ(mcp::protected_resource_metadata_path(overridden), "/custom-metadata"); + EXPECT_EQ(mcp::protected_resource_metadata_url(overridden), "https://h/custom-metadata"); +} + +TEST(ProtectedResourceMetadataTest, AnExplicitPathWithoutALeadingSlashIsRejected) { + // A relative path does not merely produce a bad path: it is concatenated straight onto the + // origin, so "https://h" + "evil" advertises the document on a host named "hevil". A typo + // that silently changes which host clients are sent to has to be reported, not published. + mcp::ProtectedResourceMetadataConfig relative; + relative.resource = "https://h/mcp"; + relative.path = "evil"; + + EXPECT_THROW(mcp::protected_resource_metadata_path(relative), std::invalid_argument); + EXPECT_THROW(mcp::protected_resource_metadata_url(relative), std::invalid_argument); +} + +TEST(ProtectedResourceMetadataTest, AnExplicitPathWithADotSegmentIsRejected) { + // The metadata URL is published to clients as the authoritative location of the document, + // and this function does not normalise. A dot segment there is never intentional + // configuration, so it is refused rather than advertised unresolved. + mcp::ProtectedResourceMetadataConfig parent; + parent.resource = "https://h/mcp"; + parent.path = "/../../x"; + EXPECT_THROW(mcp::protected_resource_metadata_path(parent), std::invalid_argument); + + mcp::ProtectedResourceMetadataConfig current; + current.resource = "https://h/mcp"; + current.path = "/a/./b"; + EXPECT_THROW(mcp::protected_resource_metadata_path(current), std::invalid_argument); + + // Only a whole segment is a dot segment. Dots inside a segment are ordinary characters, and + // rejecting those would refuse the well-known prefix this module derives itself. + mcp::ProtectedResourceMetadataConfig dotted; + dotted.resource = "https://h/mcp"; + dotted.path = "/.well-known/a..b/c.d"; + EXPECT_EQ(mcp::protected_resource_metadata_path(dotted), "/.well-known/a..b/c.d"); + EXPECT_EQ(mcp::protected_resource_metadata_url(dotted), "https://h/.well-known/a..b/c.d"); +} + +TEST(ProtectedResourceMetadataTest, DocumentUrlIgnoresQueryAndFragmentOnTheResource) { + mcp::ProtectedResourceMetadataConfig noisy; + noisy.resource = "https://h/mcp?x=1#frag"; + + EXPECT_EQ(mcp::protected_resource_metadata_path(noisy), + "/.well-known/oauth-protected-resource/mcp"); +} + +TEST(ProtectedResourceMetadataTest, DocumentOmitsEmptyLists) { + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "http://127.0.0.1:9000/mcp"; + + const auto document = nlohmann::json::parse(mcp::format_protected_resource_metadata(metadata)); + + EXPECT_EQ(document.at("resource"), "http://127.0.0.1:9000/mcp"); + EXPECT_FALSE(document.contains("authorization_servers")); + EXPECT_FALSE(document.contains("scopes_supported")); +} + +// =========================================================================== +// Server-side OAuth challenge, metadata route and unauthenticated paths +// =========================================================================== + +namespace { + +struct ChallengeProbeResult { + unsigned int status{0}; + std::string www_authenticate; + std::string content_type; + std::string body; +}; + +/// Fire one HTTP request at the transport and report the parts the challenge tests assert on. +mcp::Task probe(const asio::any_io_executor& executor, unsigned short port, + http::verb method, const std::string& target, + const std::string& body = {}, + const std::string& bearer_token = {}, + const std::string& origin = {}) { + beast::tcp_stream stream(executor); + asio::ip::tcp::resolver resolver(executor); + auto endpoints = + co_await resolver.async_resolve("127.0.0.1", std::to_string(port), asio::use_awaitable); + co_await stream.async_connect(*endpoints.begin(), asio::use_awaitable); + + http::request request{method, target, 11}; + request.set(http::field::host, "127.0.0.1"); + request.set(http::field::content_type, "application/json"); + if (!bearer_token.empty()) { + request.set(http::field::authorization, "Bearer " + bearer_token); + } + if (!origin.empty()) { + request.set(http::field::origin, origin); + } + request.body() = body; + request.prepare_payload(); + co_await http::async_write(stream, request, asio::use_awaitable); + + beast::flat_buffer response_buffer; + http::response response; + co_await http::async_read(stream, response_buffer, response, asio::use_awaitable); + + ChallengeProbeResult result; + result.status = response.result_int(); + result.www_authenticate = std::string(response[http::field::www_authenticate]); + result.content_type = std::string(response[http::field::content_type]); + result.body = response.body(); + + beast::error_code shutdown_error; + stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + co_return result; +} + +constexpr std::string_view g_notification_body = R"({"jsonrpc":"2.0","method":"notifications/x"})"; + +} // namespace + +// The challenge URL must come from `resource` alone. Behind a TLS terminator, a reverse proxy or a +// container port mapping the listener's own origin is not the one clients can reach, so a URL +// inferred from it would be unfetchable. The listener below is deliberately bound to an ephemeral +// port that has nothing to do with the advertised resource. +TEST_F(HttpTransportTest, ChallengeMetadataUrlComesFromTheResourceNotTheListener) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + const auto listener_port = server->port(); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "https://mcp.example.com/mcp"; + server->set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied = co_await probe(io_ctx_.get_executor(), listener_port, http::verb::post, "/mcp", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ( + denied.www_authenticate, + R"(Bearer resource_metadata="https://mcp.example.com/.well-known/oauth-protected-resource/mcp")"); + EXPECT_EQ(denied.www_authenticate.find("127.0.0.1"), std::string::npos) + << "the challenge leaked the listener address: " << denied.www_authenticate; + EXPECT_EQ(denied.www_authenticate.find(std::to_string(listener_port)), std::string::npos) + << "the challenge leaked the listener port: " << denied.www_authenticate; +} + +TEST_F(HttpTransportTest, UnconfiguredTransportSendsBareBearerChallenge) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(denied.www_authenticate, "Bearer"); +} + +TEST_F(HttpTransportTest, ConfiguredChallengeIsSentOnUnauthorized) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + + mcp::BearerChallengeConfig challenge; + challenge.realm = "mcp"; + challenge.scope = "mcp:read"; + challenge.resource_metadata = "http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp"; + server->set_bearer_challenge(challenge); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(denied.www_authenticate, + R"(Bearer realm="mcp", scope="mcp:read", )" + R"(resource_metadata="http://127.0.0.1:9000/.well-known/oauth-protected-resource/mcp")"); +} + +TEST_F(HttpTransportTest, ProtectedResourceMetadataIsReadableWithoutAToken) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + const auto port = server->port(); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "http://127.0.0.1:" + std::to_string(port) + "/mcp"; + metadata.authorization_servers = {"http://127.0.0.1:9000"}; + metadata.scopes_supported = {"mcp:read"}; + server->set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult document; + ChallengeProbeResult denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + document = co_await probe(io_ctx_.get_executor(), port, http::verb::get, + "/.well-known/oauth-protected-resource/mcp"); + denied = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + ASSERT_EQ(document.status, 200); + EXPECT_EQ(document.content_type, "application/json"); + const auto parsed = nlohmann::json::parse(document.body); + EXPECT_EQ(parsed.at("resource"), "http://127.0.0.1:" + std::to_string(port) + "/mcp"); + EXPECT_EQ(parsed.at("authorization_servers"), nlohmann::json::array({"http://127.0.0.1:9000"})); + EXPECT_EQ(parsed.at("scopes_supported"), nlohmann::json::array({"mcp:read"})); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(denied.www_authenticate, R"(Bearer resource_metadata="http://127.0.0.1:)" + + std::to_string(port) + + R"(/.well-known/oauth-protected-resource/mcp")"); +} + +TEST_F(HttpTransportTest, ExplicitChallengeMetadataUrlSurvivesMetadataConfiguration) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + const auto port = server->port(); + + mcp::BearerChallengeConfig challenge; + challenge.resource_metadata = "http://gateway.example/.well-known/oauth-protected-resource"; + server->set_bearer_challenge(challenge); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "http://127.0.0.1:" + std::to_string(port) + "/mcp"; + server->set_protected_resource_metadata(metadata); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult denied; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ( + denied.www_authenticate, + R"(Bearer resource_metadata="http://gateway.example/.well-known/oauth-protected-resource")"); +} + +TEST_F(HttpTransportTest, UnauthenticatedPathsAreExemptFromAuthAndExcludedFromDispatch) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + server->set_unauthenticated_paths({"/health"}); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult exempt; + ChallengeProbeResult exempt_with_query; + ChallengeProbeResult guarded; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + exempt = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/health", + std::string(g_notification_body)); + exempt_with_query = co_await probe(io_ctx_.get_executor(), port, http::verb::post, + "/health?probe=1", std::string(g_notification_body)); + guarded = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/healthy", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + // Exempt from the bearer check, and excluded from MCP dispatch: nothing claimed the path, so + // it is 404 rather than an unauthenticated 202. + EXPECT_EQ(exempt.status, 404); + EXPECT_EQ(exempt_with_query.status, 404); + EXPECT_EQ(guarded.status, 401); +} + +// The catastrophic misconfiguration: exempting the path MCP is served on. It must fail loudly. +TEST_F(HttpTransportTest, ExemptingTheMcpPathRefusesToServeMcpUnauthenticated) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + server->set_unauthenticated_paths({"/mcp"}); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult dispatched; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + dispatched = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(dispatched.status, 404); + EXPECT_NE(dispatched.status, 202) << "an exempt path must never reach unauthenticated dispatch"; +} + +TEST_F(HttpTransportTest, AnExemptPathTheMetadataRouteClaimsIsStillServed) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + const auto port = server->port(); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "https://mcp.example.com/mcp"; + server->set_protected_resource_metadata(metadata); + server->set_unauthenticated_paths({"/.well-known/oauth-protected-resource/mcp"}); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult document; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + document = co_await probe(io_ctx_.get_executor(), port, http::verb::get, + "/.well-known/oauth-protected-resource/mcp"); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + ASSERT_EQ(document.status, 200); + EXPECT_EQ(nlohmann::json::parse(document.body).at("resource"), "https://mcp.example.com/mcp"); +} + +TEST_F(HttpTransportTest, AsyncBearerValidatorDecidesWithoutBlocking) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + std::atomic validator_calls{0}; + server->set_async_bearer_token_validator([&validator_calls](std::string token) -> mcp::Task { + validator_calls.fetch_add(1, std::memory_order_relaxed); + // Suspend, the way a real introspection call would, before deciding. + asio::steady_timer timer(co_await asio::this_coro::executor); + timer.expires_after(std::chrono::milliseconds(1)); + co_await timer.async_wait(asio::use_awaitable); + co_return token == "good"; + }); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult denied; + ChallengeProbeResult accepted; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + denied = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body), "bad"); + accepted = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body), "good"); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(denied.status, 401); + EXPECT_EQ(denied.www_authenticate, "Bearer"); + EXPECT_EQ(accepted.status, 202); + EXPECT_EQ(validator_calls.load(std::memory_order_relaxed), 2); +} + +TEST_F(HttpTransportTest, OnlyOneBearerValidatorMayBeInstalled) { + mcp::HttpServerTransport sync_first(io_ctx_.get_executor(), "127.0.0.1", 0); + sync_first.set_bearer_token_validator([](std::string_view) { return true; }); + EXPECT_THROW(sync_first.set_async_bearer_token_validator( + [](std::string) -> mcp::Task { co_return true; }), + std::logic_error); + sync_first.close(); + + mcp::HttpServerTransport async_first(io_ctx_.get_executor(), "127.0.0.1", 0); + async_first.set_async_bearer_token_validator( + [](std::string) -> mcp::Task { co_return true; }); + EXPECT_THROW(async_first.set_bearer_token_validator([](std::string_view) { return true; }), + std::logic_error); + async_first.close(); +} + +TEST_F(HttpTransportTest, RequestBodyBeyondTheLimitIsRejected) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + unsigned int status = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + // Announcing the length is enough: the parser rejects the request before the body is + // sent, which is the point of the limit. + const std::string header = + "POST /mcp HTTP/1.1\r\nHost: 127.0.0.1\r\nContent-Type: application/json\r\n" + "Content-Length: " + + std::to_string(mcp::constants::g_default_max_request_body_bytes + 1) + "\r\n\r\n"; + + beast::tcp_stream stream(io_ctx_.get_executor()); + asio::ip::tcp::resolver resolver(io_ctx_.get_executor()); + auto endpoints = + co_await resolver.async_resolve("127.0.0.1", std::to_string(port), asio::use_awaitable); + co_await stream.async_connect(*endpoints.begin(), asio::use_awaitable); + co_await asio::async_write(stream, asio::buffer(header), asio::use_awaitable); + + beast::flat_buffer response_buffer; + http::response response; + co_await http::async_read(stream, response_buffer, response, asio::use_awaitable); + status = response.result_int(); + + beast::error_code shutdown_error; + stream.socket().shutdown(asio::ip::tcp::socket::shutdown_both, shutdown_error); + }, + [server](const std::exception_ptr&) { server->close(); }); + io_ctx_.run(); + + EXPECT_EQ(status, 413); +} + +TEST_F(HttpTransportTest, ProtectedResourceMetadataRequiresAnAbsoluteResourceUrl) { + mcp::HttpServerTransport server(io_ctx_.get_executor(), "127.0.0.1", 0); + + EXPECT_THROW(server.set_protected_resource_metadata({}), std::invalid_argument); + + mcp::ProtectedResourceMetadataConfig relative; + relative.resource = "/mcp"; + EXPECT_THROW(server.set_protected_resource_metadata(relative), std::invalid_argument); + + server.close(); +} + +// =========================================================================== +// Origin checking (DNS rebinding protection) on HttpServerTransport +// +// The conformance runner's dns-rebinding-protection scenario drives +// StreamableHttpSessionManager, which is what conformance/everything_server.cpp binds. Nothing +// exercises the same defence on HttpServerTransport, and the two transports do not answer a +// disallowed origin the same way, so the difference is pinned here rather than assumed away. +// =========================================================================== + +namespace { + +/// Post one MCP notification carrying `origin` and report the status the transport answered with. +mcp::Task origin_probe_status(const asio::any_io_executor& executor, unsigned short port, + const std::string& origin) { + const auto result = co_await probe(executor, port, http::verb::post, "/mcp", + std::string(g_notification_body), {}, origin); + co_return result.status; +} + +} // namespace + +// A transport that was never told which origins may reach it refuses every request that names +// one. The check fails closed: an empty allow-list is a deny-all, not an "unconfigured, so skip +// it", which is what makes the default safe against a rebound DNS name. +TEST_F(HttpTransportTest, UnconfiguredTransportRefusesAnyOriginHeader) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + unsigned int status = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + status = + co_await origin_probe_status(io_ctx_.get_executor(), port, "https://untrusted.example"); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(status, 403u); +} + +// Only a request that names an origin is subject to the check. A browser sends Origin; a CLI, a +// proxy health probe and the SDK's own HttpClientTransport do not, and refusing those would make +// the safe default unusable for every non-browser client. +TEST_F(HttpTransportTest, ARequestWithNoOriginHeaderSkipsTheOriginCheck) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult result; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + result = co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body)); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(result.status, 202u); +} + +TEST_F(HttpTransportTest, NamedOriginIsAdmittedAndAnUnlistedOneIsRefused) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_allowed_origins({"https://trusted.example", "https://also-trusted.example"}); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + unsigned int first_allowed = 0; + unsigned int second_allowed = 0; + unsigned int refused = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + first_allowed = + co_await origin_probe_status(io_ctx_.get_executor(), port, "https://trusted.example"); + second_allowed = co_await origin_probe_status(io_ctx_.get_executor(), port, + "https://also-trusted.example"); + refused = + co_await origin_probe_status(io_ctx_.get_executor(), port, "https://untrusted.example"); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(first_allowed, 202u); + EXPECT_EQ(second_allowed, 202u); + EXPECT_EQ(refused, 403u); +} + +// The allow-list is a set of exact strings, not the canonicalizing comparison the client-side +// MetadataFetchPolicy performs on the origins it will fetch from. An explicit default port, a +// different scheme or host case and a trailing slash all denote the same web origin, and all of +// them are refused here. A deployment that wants them admitted must list every spelling. +TEST_F(HttpTransportTest, AllowedOriginComparisonIsExactRatherThanCanonicalizing) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_allowed_origins({"https://trusted.example"}); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + const std::vector spellings = {"https://trusted.example:443", + "HTTPS://trusted.example", "https://Trusted.example", + "https://trusted.example/"}; + std::vector statuses; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + for (const auto& origin : spellings) { + statuses.push_back(co_await origin_probe_status(io_ctx_.get_executor(), port, origin)); + } + server->close(); + }, + asio::detached); + io_ctx_.run(); + + ASSERT_EQ(statuses.size(), spellings.size()); + for (std::size_t index = 0; index < spellings.size(); ++index) { + SCOPED_TRACE(spellings[index]); + EXPECT_EQ(statuses[index], 403u); + } +} + +// The documented escape hatch for a deployment that fronts the transport with its own origin +// policy. Nothing else has to be configured for it to take effect. +TEST_F(HttpTransportTest, AllowAllOriginsAdmitsAnOriginThatWasNeverListed) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_allow_all_origins(true); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + unsigned int status = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + status = + co_await origin_probe_status(io_ctx_.get_executor(), port, "https://untrusted.example"); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(status, 202u); +} + +// set_allowed_origins() names the origins that may connect, so it also revokes a blanket +// allowance granted earlier. Leaving allow-all in force behind an allow-list would make the +// narrower call silently do nothing. +TEST_F(HttpTransportTest, NamingAllowedOriginsRevokesAnEarlierAllowAll) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_allow_all_origins(true); + server->set_allowed_origins({"https://trusted.example"}); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + unsigned int allowed = 0; + unsigned int refused = 0; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + allowed = + co_await origin_probe_status(io_ctx_.get_executor(), port, "https://trusted.example"); + refused = + co_await origin_probe_status(io_ctx_.get_executor(), port, "https://untrusted.example"); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(allowed, 202u); + EXPECT_EQ(refused, 403u); +} + +// Where the origin check sits in this transport's pipeline: not once up front as in +// StreamableHttpSessionManager, but on each route that needs it -- before the bearer check on the MCP +// path, and inside the RFC 9728 metadata route. A rebinding attempt is refused 403 on both, and the +// MCP refusal carries no WWW-Authenticate challenge. +// +// The unauthenticated-path list is the one route the check does not reach: handle_request() answers +// an exempt path 404 before dispatch, where the session manager returns 403 for the same request. See +// the companion test in transport_http_session_manager_test.cpp. +TEST_F(HttpTransportTest, TheOriginCheckGuardsMcpDispatchAndTheMetadataRouteButNotExemptPaths) { + auto server = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", 0); + server->set_allowed_origins({"https://trusted.example"}); + server->set_bearer_token_validator([](std::string_view token) { return token == "good"; }); + server->set_unauthenticated_paths({"/health"}); + + mcp::ProtectedResourceMetadataConfig metadata; + metadata.resource = "https://mcp.example.com/mcp"; + server->set_protected_resource_metadata(metadata); + const auto port = server->port(); + + asio::co_spawn(io_ctx_, server->listen(), asio::detached); + + ChallengeProbeResult mcp_result; + ChallengeProbeResult metadata_result; + ChallengeProbeResult exempt_result; + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + mcp_result = + co_await probe(io_ctx_.get_executor(), port, http::verb::post, "/mcp", + std::string(g_notification_body), {}, "https://untrusted.example"); + metadata_result = co_await probe(io_ctx_.get_executor(), port, http::verb::get, + "/.well-known/oauth-protected-resource/mcp", {}, {}, + "https://untrusted.example"); + exempt_result = co_await probe(io_ctx_.get_executor(), port, http::verb::get, "/health", {}, + {}, "https://untrusted.example"); + server->close(); + }, + asio::detached); + io_ctx_.run(); + + EXPECT_EQ(mcp_result.status, 403u); + EXPECT_TRUE(mcp_result.www_authenticate.empty()); + + EXPECT_EQ(metadata_result.status, 403u); + + EXPECT_EQ(exempt_result.status, 404u); +} + +namespace { + +struct CloseAfterResponseAttempt { + bool session_started{false}; + bool request_held{false}; + bool close_ran{false}; + bool response_sent{false}; + bool write_finished{false}; + bool write_succeeded{false}; + bool delete_received{false}; + int deletes{0}; + int accepted{0}; +}; + +/// One run of the scenario below. `settle` is how long the io thread stays held after the peer has +/// written the response, before close() is called. +CloseAfterResponseAttempt close_after_the_response_arrived(std::chrono::milliseconds settle) { + CloseAfterResponseAttempt attempt; + SessionPeer peer; + asio::io_context io_ctx; + auto transport = std::make_shared( + io_ctx.get_executor(), "http://127.0.0.1:" + std::to_string(peer.port()) + "/mcp"); + + std::exception_ptr initialize_failure; + std::exception_ptr write_failure; + std::promise initialize_done; + std::promise write_done; + auto initialized = initialize_done.get_future(); + auto write_finished = write_done.get_future(); + asio::co_spawn( + io_ctx, + [&]() -> mcp::Task { + try { + co_await transport->write_message(initialize_request_text()); + co_await transport->read_message(); + } catch (...) { + initialize_failure = std::current_exception(); + } + initialize_done.set_value(); + if (initialize_failure) { + co_return; + } + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":2,"method":"ping"})"); + } catch (...) { + write_failure = std::current_exception(); + } + write_done.set_value(); + }, + asio::detached); + + auto work = asio::make_work_guard(io_ctx); + std::thread runner([&io_ctx]() { io_ctx.run(); }); + + const auto limit = std::chrono::seconds(10); + attempt.session_started = initialized.wait_for(limit) == std::future_status::ready; + attempt.request_held = attempt.session_started && peer.wait_for_held_request(limit); + + // The peer has the whole request, so the handler that wrote it has run. Two hops through the + // io_context later the write has also resumed from it and is parked in its response read. + std::promise parked; + auto write_parked = parked.get_future(); + std::promise closed; + auto close_queued = closed.get_future(); + bool response_sent = false; + bool close_posted = false; + if (attempt.request_held) { + asio::post(io_ctx, [&]() { asio::post(io_ctx, [&]() { parked.set_value(); }); }); + if (write_parked.wait_for(limit) == std::future_status::ready) { + close_posted = true; + asio::post(io_ctx, [&]() { + // The io thread is held here: nothing of the transport runs. + peer.release_held_response(); + response_sent = peer.wait_for_held_response_sent(limit); + std::this_thread::sleep_for(settle); + std::thread closer([&]() { transport->close(); }); + closer.join(); + closed.set_value(); + }); + } + } + attempt.close_ran = + close_posted && close_queued.wait_for(std::chrono::seconds(30)) == std::future_status::ready; + attempt.write_finished = + attempt.close_ran && write_finished.wait_for(limit) == std::future_status::ready; + attempt.delete_received = attempt.write_finished && peer.wait_for_deletes(1, limit); + + // Everything below reads state the runner thread wrote, so it stops first. + io_ctx.stop(); + runner.join(); + + attempt.response_sent = response_sent; + attempt.write_succeeded = attempt.write_finished && !initialize_failure && write_failure == nullptr; + attempt.deletes = peer.deletes(); + attempt.accepted = peer.accepted(); + return attempt; +} + +} // namespace + +// close() called from another thread when the response to the write in flight has already arrived +// but the io thread has not read it yet. A handler holds the single io thread while the peer sends +// the response and a second thread calls close(), so close()'s work is queued ahead of the read's +// completion: it closes the socket of a write that then finishes successfully. The session DELETE +// that follows must not be attempted on that closed socket, where it would fail and be dropped. +// +// The test cannot observe when the peer's bytes become readable on the transport's socket. If close() +// runs before they are, the write fails instead and the run proves nothing, so the scenario is run +// again holding the io thread longer, until the write has succeeded. +TEST_F(HttpTransportTest, CloseAfterTheResponseArrivedStillSendsTheSessionDelete) { + CloseAfterResponseAttempt attempt; + auto settle = std::chrono::milliseconds(0); + for (int run = 0; run < 6 && !attempt.write_succeeded; ++run) { + attempt = close_after_the_response_arrived(settle); + ASSERT_TRUE(attempt.session_started); + ASSERT_TRUE(attempt.request_held) << "the second write never reached the peer"; + ASSERT_TRUE(attempt.close_ran); + ASSERT_TRUE(attempt.response_sent); + ASSERT_TRUE(attempt.write_finished); + settle = settle.count() == 0 ? std::chrono::milliseconds(50) : settle * 4; + } + + ASSERT_TRUE(attempt.write_succeeded) + << "the write did not complete in any run, including the ones that held the io thread for " + "seconds after the peer had sent the response: suspect the transport's read or its " + "close(), not the timing of this test"; + EXPECT_TRUE(attempt.delete_received) << "close() dropped the session DELETE"; + EXPECT_EQ(attempt.deletes, 1); + EXPECT_EQ(attempt.accepted, 2) << "the DELETE is expected on a connection of its own"; +} + +// close() called from another thread while the write is between two socket operations: connected, +// nothing pending. The bearer provider runs on the transport's strand right before the request is +// written, so calling close() from a second thread inside it puts close()'s work behind exactly +// that point. A close() that only cancels what is pending finds nothing, the request is written, +// and the write then waits on the stalling server until the HTTP timeout. The close lands on no +// pending operation, so the socket layer reports a closed socket rather than a cancelled operation; +// the write reports the cancellation all the same. +TEST_F(HttpTransportTest, CloseBetweenTwoSocketOperationsEndsTheWrite) { + StallingServer stalling(io_ctx_); + stalling.accept_and_stall(); + + auto transport = std::make_shared( + io_ctx_.get_executor(), "http://127.0.0.1:" + std::to_string(stalling.port()) + "/mcp"); + std::weak_ptr weak_transport = transport; + transport->set_bearer_token_provider([weak_transport]() { + if (auto locked = weak_transport.lock()) { + std::thread closer([locked]() { locked->close(); }); + closer.join(); + } + return std::string("token"); + }); + + std::exception_ptr failure; + std::promise write_done; + auto write_finished = write_done.get_future(); + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"ping"})"); + } catch (...) { + failure = std::current_exception(); + } + write_done.set_value(); + }, + asio::detached); + + std::thread runner([this]() { io_ctx_.run(); }); + + const bool finished = + write_finished.wait_for(std::chrono::seconds(10)) == std::future_status::ready; + + // Everything below reads state the runner thread wrote, so it stops first. + io_ctx_.stop(); + runner.join(); + + ASSERT_TRUE(finished) << "close() was lost: the write went on to wait for a response after " + "the transport closed and is still running"; + ASSERT_NE(failure, nullptr) << "a write cut short by close() must report an error"; + try { + std::rethrow_exception(failure); + } catch (const boost::system::system_error& error) { + EXPECT_EQ(error.code(), asio::error::operation_aborted) + << "the write ended with: " << error.code().message(); + } catch (const std::exception& other) { + ADD_FAILURE() << "the write ended with: " << other.what(); + } +} + +#ifdef __linux__ + +// close() called from a thread that does not run the io_context, while the write is suspended in +// its resolve. Asio's resolver thread checks its cancel token once, before it calls getaddrinfo(), +// so a close() that arrives while that call is in progress finds a resolve it can no longer cancel +// and no socket to cancel either. The gate holds the lookup inside getaddrinfo() until close() has +// done its cancelling, so the write resumes with usable addresses on a transport that is already +// closed. It must stop there. Without the closed check in ensure_connected() it connects to the +// stalling server and stays blocked on it until the HTTP timeout. +TEST_F(HttpTransportTest, CloseWhileTheResolverIsPastItsCancelCheckOpensNoConnection) { + StallingServer stalling(io_ctx_); + stalling.accept_and_stall(); + + auto transport = std::make_shared( + io_ctx_.get_executor(), "http://127.0.0.1:" + std::to_string(stalling.port()) + "/mcp"); + + resolve_gate().arm(stalling.port()); + + std::exception_ptr failure; + std::promise write_done; + auto write_finished = write_done.get_future(); + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"ping"})"); + } catch (...) { + failure = std::current_exception(); + } + write_done.set_value(); + }, + asio::detached); + + std::thread runner([this]() { io_ctx_.run(); }); + + const auto limit = std::chrono::seconds(10); + const bool lookup_held = resolve_gate().wait_until_entered(limit); + + // close() spawns its cancelling onto the transport's strand, which is one hop through the + // io_context. Two further hops through the same queue cannot complete before it has. + std::promise cancel_done; + auto cancel_finished = cancel_done.get_future(); + bool cancel_ran = false; + if (lookup_held) { + transport->close(); + asio::post(io_ctx_, [&]() { asio::post(io_ctx_, [&]() { cancel_done.set_value(); }); }); + cancel_ran = cancel_finished.wait_for(limit) == std::future_status::ready; + } + + resolve_gate().release(); + const bool finished = write_finished.wait_for(limit) == std::future_status::ready; + + // Everything below reads state the runner thread wrote, so it stops first. + io_ctx_.stop(); + runner.join(); + + ASSERT_TRUE(lookup_held) << "the write never reached getaddrinfo()"; + ASSERT_TRUE(cancel_ran) << "the cancelling spawned by close() never ran on the io_context"; + EXPECT_TRUE(finished) << "close() was lost: the write resumed from the lookup after the " + "transport closed and is still running"; + EXPECT_EQ(stalling.accepted(), 0) << "the write connected after the transport had closed"; + if (finished) { + EXPECT_NE(failure, nullptr) << "a write cut short by close() must report an error"; + } +} + +#endif // __linux__ diff --git a/test/transport/transport_memory_test.cpp b/test/transport/transport_memory_test.cpp index c52ba5e..adcd801 100644 --- a/test/transport/transport_memory_test.cpp +++ b/test/transport/transport_memory_test.cpp @@ -4,12 +4,22 @@ #include #include +#include #include #include +#include + +#include +#include +#include +#include #include +#include #include #include +#include #include +#include class MemoryTransportTest : public ::testing::Test { protected: @@ -217,3 +227,324 @@ TEST_F(MemoryTransportTest, WriteAfterCloseThrows) { EXPECT_TRUE(exception_thrown); } + +TEST_F(MemoryTransportTest, WriteOwnsStringViewBeforeTaskIsAwaited) { + auto [transport_a, transport_b] = mcp::create_memory_transport_pair(io_ctx_.get_executor()); + + std::string source = "original message"; + std::string received_message; + + boost::asio::co_spawn( + io_ctx_, + [ta = std::move(transport_a), tb = std::move(transport_b), &source, + &received_message]() mutable -> mcp::Task { + auto pending_write = ta->write_message(source); + source.assign("changed after write_message returned"); + + co_await std::move(pending_write); + received_message = co_await tb->read_message(); + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(received_message, "original message"); +} + +TEST_F(MemoryTransportTest, PendingReadRetainsEndpointAfterWrapperDestruction) { + auto [transport_a, transport_b] = mcp::create_memory_transport_pair(io_ctx_.get_executor()); + + std::string received_message; + + boost::asio::co_spawn( + io_ctx_, + [ta = std::move(transport_a), tb = std::move(transport_b), + &received_message]() mutable -> mcp::Task { + auto pending_read = tb->read_message(); + tb.reset(); + + co_await ta->write_message("message for retained endpoint"); + received_message = co_await std::move(pending_read); + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(received_message, "message for retained endpoint"); +} + +TEST_F(MemoryTransportTest, DestroyingIdleEndpointWakesPeerRead) { + auto [transport_a, transport_b] = mcp::create_memory_transport_pair(io_ctx_.get_executor()); + + bool peer_closed = false; + boost::asio::co_spawn( + io_ctx_, + [transport_a = std::move(transport_a), transport_b = std::move(transport_b), + &peer_closed]() mutable -> mcp::Task { + auto pending_read = transport_b->read_message(); + transport_a.reset(); + + try { + static_cast(co_await std::move(pending_read)); + } catch (const std::runtime_error& error) { + peer_closed = true; + EXPECT_STREQ(error.what(), "transport closed"); + } + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_TRUE(peer_closed); +} + +TEST_F(MemoryTransportTest, JsonMessagesRoundTripWithoutSerializationAtTheWriter) { + auto [base_a, base_b] = mcp::create_memory_transport_pair(io_ctx_.get_executor()); + auto transport_a = std::dynamic_pointer_cast(base_a); + auto transport_b = std::dynamic_pointer_cast(base_b); + ASSERT_NE(transport_a, nullptr); + ASSERT_NE(transport_b, nullptr); + + const nlohmann::json expected = {{"jsonrpc", "2.0"}, {"id", 7}, {"result", {}}}; + nlohmann::json received; + + boost::asio::co_spawn( + io_ctx_, + [transport_a = std::move(transport_a), transport_b = std::move(transport_b), expected, + &received]() mutable -> mcp::Task { + co_await transport_a->write_json(expected); + received = co_await transport_b->read_json(); + }, + boost::asio::detached); + + io_ctx_.run(); + + EXPECT_EQ(received, expected); +} + +TEST_F(MemoryTransportTest, ConcurrentCloseWakesPendingReadExactlyOnce) { + auto [transport_a, transport_b] = mcp::create_memory_transport_pair(io_ctx_.get_executor()); + + int close_errors = 0; + boost::asio::co_spawn( + io_ctx_, + [transport_a, &close_errors]() -> mcp::Task { + try { + (void)co_await transport_a->read_message(); + } catch (const std::runtime_error& error) { + EXPECT_STREQ(error.what(), "transport closed"); + ++close_errors; + } + }, + boost::asio::detached); + + io_ctx_.poll(); + + std::vector closers; + closers.reserve(8); + for (int index = 0; index < 8; ++index) { + auto transport = index % 2 == 0 ? transport_a : transport_b; + closers.emplace_back([transport]() { transport->close(); }); + } + for (auto& closer : closers) { + closer.join(); + } + + io_ctx_.restart(); + io_ctx_.run(); + + EXPECT_EQ(close_errors, 1); +} + +TEST_F(MemoryTransportTest, ConcurrentReadWriteIsSafeAcrossThreadPool) { + constexpr int round_count = 4; + constexpr int message_count = 256; + + for (int round = 0; round < round_count; ++round) { + boost::asio::io_context io_context; + auto [transport_a, transport_b] = mcp::create_memory_transport_pair(io_context.get_executor()); + + std::atomic write_completions{0}; + std::atomic failures{0}; + std::vector received; + received.reserve(message_count); + + boost::asio::co_spawn( + io_context, + [transport_b, &received]() -> mcp::Task { + for (int index = 0; index < message_count; ++index) { + received.push_back(co_await transport_b->read_message()); + } + }, + [&failures](std::exception_ptr error) { + if (error) { + failures.fetch_add(1, std::memory_order_relaxed); + } + }); + + std::vector expected; + expected.reserve(message_count); + for (int index = 0; index < message_count; ++index) { + auto message = "round-" + std::to_string(round) + "-message-" + std::to_string(index); + expected.push_back(message); + boost::asio::co_spawn( + io_context, + [transport_a, message = std::move(message)]() -> mcp::Task { + co_await transport_a->write_message(message); + }, + [&write_completions, &failures](std::exception_ptr error) { + if (error) { + failures.fetch_add(1, std::memory_order_relaxed); + } + write_completions.fetch_add(1, std::memory_order_release); + }); + } + + std::vector workers; + workers.reserve(4); + for (int index = 0; index < 4; ++index) { + workers.emplace_back([&io_context]() { io_context.run(); }); + } + for (auto& worker : workers) { + worker.join(); + } + + std::sort(received.begin(), received.end()); + std::sort(expected.begin(), expected.end()); + EXPECT_EQ(failures.load(std::memory_order_acquire), 0) << "round " << round; + EXPECT_EQ(write_completions.load(std::memory_order_acquire), message_count) + << "round " << round; + EXPECT_EQ(received, expected) << "round " << round; + } +} + +TEST_F(MemoryTransportTest, MultiplePendingReadersReceiveDistinctMessages) { + boost::asio::io_context io_context; + auto [transport_a, transport_b] = mcp::create_memory_transport_pair(io_context.get_executor()); + + constexpr int reader_count = 64; + std::atomic failures{0}; + std::vector received(reader_count); + + for (int index = 0; index < reader_count; ++index) { + boost::asio::co_spawn( + io_context, + [transport_b, &received, index]() -> mcp::Task { + received[index] = co_await transport_b->read_message(); + }, + [&failures](std::exception_ptr error) { + if (error) { + failures.fetch_add(1, std::memory_order_relaxed); + } + }); + } + + // Install every pending read before any message is delivered. This specifically + // exercises independent waiters rather than the already-queued fast path. + io_context.poll(); + io_context.restart(); + + std::vector expected; + expected.reserve(reader_count); + for (int index = 0; index < reader_count; ++index) { + auto message = "pending-reader-message-" + std::to_string(index); + expected.push_back(message); + boost::asio::co_spawn( + io_context, + [transport_a, message = std::move(message)]() -> mcp::Task { + co_await transport_a->write_message(message); + }, + [&failures](std::exception_ptr error) { + if (error) { + failures.fetch_add(1, std::memory_order_relaxed); + } + }); + } + + std::vector workers; + workers.reserve(4); + for (int index = 0; index < 4; ++index) { + workers.emplace_back([&io_context]() { io_context.run(); }); + } + for (auto& worker : workers) { + worker.join(); + } + + std::sort(received.begin(), received.end()); + std::sort(expected.begin(), expected.end()); + EXPECT_EQ(failures.load(std::memory_order_acquire), 0); + EXPECT_EQ(received, expected); +} + +TEST_F(MemoryTransportTest, ConcurrentCloseWakesReadersOnRunningThreadPool) { + boost::asio::io_context io_context; + auto work_guard = boost::asio::make_work_guard(io_context); + auto [transport_a, transport_b] = mcp::create_memory_transport_pair(io_context.get_executor()); + + constexpr int reader_count = 32; + std::atomic started{0}; + std::atomic close_errors{0}; + std::atomic unexpected_completions{0}; + std::mutex started_mutex; + std::condition_variable started_condition; + + for (int index = 0; index < reader_count; ++index) { + boost::asio::co_spawn( + io_context, + [transport_a, &started, &close_errors, &unexpected_completions, &started_mutex, + &started_condition]() -> mcp::Task { + { + // The increment must happen under the mutex: notifying while the + // waiter is between its predicate check and blocking would + // otherwise be lost and the main thread would wait forever. + std::lock_guard lock(started_mutex); + started.fetch_add(1, std::memory_order_release); + } + started_condition.notify_one(); + try { + (void)co_await transport_a->read_message(); + unexpected_completions.fetch_add(1, std::memory_order_relaxed); + } catch (const std::runtime_error& error) { + if (std::string_view(error.what()) == "transport closed") { + close_errors.fetch_add(1, std::memory_order_relaxed); + } else { + unexpected_completions.fetch_add(1, std::memory_order_relaxed); + } + } + }, + boost::asio::detached); + } + + std::vector workers; + workers.reserve(4); + for (int index = 0; index < 4; ++index) { + workers.emplace_back([&io_context]() { io_context.run(); }); + } + + bool all_readers_started = false; + { + std::unique_lock lock(started_mutex); + all_readers_started = started_condition.wait_for(lock, std::chrono::seconds(30), [&started]() { + return started.load(std::memory_order_acquire) == reader_count; + }); + } + + std::vector closers; + closers.reserve(8); + for (int index = 0; index < 8; ++index) { + auto transport = index % 2 == 0 ? transport_a : transport_b; + closers.emplace_back([transport]() { transport->close(); }); + } + for (auto& closer : closers) { + closer.join(); + } + + work_guard.reset(); + for (auto& worker : workers) { + worker.join(); + } + + ASSERT_TRUE(all_readers_started) << "watchdog: readers did not all start within 30s"; + EXPECT_EQ(close_errors.load(std::memory_order_acquire), reader_count); + EXPECT_EQ(unexpected_completions.load(std::memory_order_acquire), 0); +} diff --git a/test/transport/transport_stdio_test.cpp b/test/transport/transport_stdio_test.cpp index c0dc01f..c61d09f 100644 --- a/test/transport/transport_stdio_test.cpp +++ b/test/transport/transport_stdio_test.cpp @@ -2,12 +2,182 @@ #include +#include #include #include #include #include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include #include +#include +#include #include +#include +#include +#include + +namespace { + +class BlockingStreambuf : public std::streambuf { + public: + void feed(std::string data) { + { + std::lock_guard lock(mutex_); + buffer_ += std::move(data); + } + condition_.notify_all(); + } + + void close() { + { + std::lock_guard lock(mutex_); + closed_ = true; + } + condition_.notify_all(); + } + + protected: + int_type underflow() override { + std::unique_lock lock(mutex_); + condition_.wait(lock, [this] { return position_ < buffer_.size() || closed_; }); + if (position_ == buffer_.size()) { + return traits_type::eof(); + } + return traits_type::to_int_type(buffer_[position_]); + } + + int_type uflow() override { + std::unique_lock lock(mutex_); + condition_.wait(lock, [this] { return position_ < buffer_.size() || closed_; }); + if (position_ == buffer_.size()) { + return traits_type::eof(); + } + return traits_type::to_int_type(buffer_[position_++]); + } + + private: + std::mutex mutex_; + std::condition_variable condition_; + std::string buffer_; + std::size_t position_{0}; + bool closed_{false}; +}; + +class TrackingOutputBuffer : public std::streambuf { + public: + bool observed_overlap() const noexcept { return observed_overlap_.load(std::memory_order_acquire); } + + std::string str() const { + std::lock_guard lock(data_mutex_); + return data_; + } + + protected: + std::streamsize xsputn(const char_type* data, std::streamsize size) override { + Activity activity(*this); + std::this_thread::sleep_for(std::chrono::microseconds(200)); + std::lock_guard lock(data_mutex_); + data_.append(data, static_cast(size)); + return size; + } + + int_type overflow(int_type value) override { + if (traits_type::eq_int_type(value, traits_type::eof())) { + return traits_type::not_eof(value); + } + + Activity activity(*this); + std::this_thread::sleep_for(std::chrono::microseconds(200)); + std::lock_guard lock(data_mutex_); + data_.push_back(traits_type::to_char_type(value)); + return value; + } + + int sync() override { + Activity activity(*this); + std::this_thread::sleep_for(std::chrono::microseconds(200)); + return 0; + } + + private: + class Activity { + public: + explicit Activity(TrackingOutputBuffer& owner) : owner_(owner) { + if (owner_.active_calls_.fetch_add(1, std::memory_order_acq_rel) != 0) { + owner_.observed_overlap_.store(true, std::memory_order_release); + } + } + + ~Activity() { owner_.active_calls_.fetch_sub(1, std::memory_order_acq_rel); } + + private: + TrackingOutputBuffer& owner_; + }; + + mutable std::mutex data_mutex_; + std::string data_; + std::atomic active_calls_{0}; + std::atomic observed_overlap_{false}; +}; + +class BlockingOutputBuffer : public std::streambuf { + public: + bool wait_until_write_starts(std::chrono::milliseconds timeout) { + std::unique_lock lock(mutex_); + return condition_.wait_for(lock, timeout, [this] { return write_started_; }); + } + + void release_write() { + { + std::lock_guard lock(mutex_); + write_released_ = true; + } + condition_.notify_all(); + } + + std::string str() const { + std::lock_guard lock(mutex_); + return data_; + } + + protected: + std::streamsize xsputn(const char_type* data, std::streamsize size) override { + std::unique_lock lock(mutex_); + write_started_ = true; + condition_.notify_all(); + condition_.wait(lock, [this] { return write_released_; }); + data_.append(data, static_cast(size)); + return size; + } + + int_type overflow(int_type value) override { + if (traits_type::eq_int_type(value, traits_type::eof())) { + return traits_type::not_eof(value); + } + std::lock_guard lock(mutex_); + data_.push_back(traits_type::to_char_type(value)); + return value; + } + + private: + mutable std::mutex mutex_; + std::condition_variable condition_; + std::string data_; + bool write_started_{false}; + bool write_released_{false}; +}; + +} // namespace class StdioTransportTest : public ::testing::Test { protected: @@ -84,6 +254,174 @@ TEST_F(StdioTransportTest, WriteMultipleMessages) { EXPECT_EQ(output.str(), "line1\nline2\n"); } +TEST_F(StdioTransportTest, ConcurrentWritesOwnPayloadAndRemainSerialized) { + constexpr std::size_t write_count = 64; + constexpr std::size_t thread_count = 4; + + std::istringstream input; + TrackingOutputBuffer output_buffer; + std::ostream output(&output_buffer); + mcp::StdioTransport transport(io_ctx_.get_executor(), input, output); + + std::vector expected; + expected.reserve(write_count); + for (std::size_t index = 0; index < write_count; ++index) { + expected.push_back("message-" + std::to_string(index) + ":" + + std::string(2048, static_cast('a' + index % 26))); + } + + std::vector> completions(write_count); + std::vector submitters; + for (std::size_t thread_index = 0; thread_index < thread_count; ++thread_index) { + submitters.emplace_back([&, thread_index]() { + for (std::size_t index = thread_index; index < write_count; index += thread_count) { + auto source = expected[index]; + auto write = transport.write_message(source); + source.assign("mutated after write_message returned"); + completions[index] = + boost::asio::co_spawn(io_ctx_, std::move(write), boost::asio::use_future); + } + }); + } + for (auto& submitter : submitters) { + submitter.join(); + } + + std::vector runners; + for (std::size_t index = 0; index < thread_count; ++index) { + runners.emplace_back([this]() { io_ctx_.run(); }); + } + for (auto& runner : runners) { + runner.join(); + } + for (auto& completion : completions) { + EXPECT_NO_THROW(completion.get()); + } + + EXPECT_FALSE(output_buffer.observed_overlap()); + std::unordered_set actual; + std::istringstream lines(output_buffer.str()); + for (std::string line; std::getline(lines, line);) { + actual.insert(std::move(line)); + } + EXPECT_EQ(actual.size(), write_count); + for (const auto& message : expected) { + EXPECT_EQ(actual.count(message), 1); + } +} + +TEST_F(StdioTransportTest, RejectsASecondOutstandingRead) { + BlockingStreambuf input_buffer; + std::istream input(&input_buffer); + std::ostringstream output; + mcp::StdioTransport transport(io_ctx_.get_executor(), input, output); + + auto first = boost::asio::co_spawn(io_ctx_, transport.read_message(), boost::asio::use_future); + auto second = boost::asio::co_spawn(io_ctx_, transport.read_message(), boost::asio::use_future); + + std::thread first_runner([this]() { io_ctx_.run(); }); + std::thread second_runner([this]() { io_ctx_.run(); }); + + const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(30); + bool one_completed = false; + while (std::chrono::steady_clock::now() < deadline) { + one_completed = first.wait_for(std::chrono::milliseconds(0)) == std::future_status::ready || + second.wait_for(std::chrono::milliseconds(0)) == std::future_status::ready; + if (one_completed) { + break; + } + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + + input_buffer.feed("accepted\n"); + input_buffer.close(); + first_runner.join(); + second_runner.join(); + + EXPECT_TRUE(one_completed); + unsigned int accepted = 0; + unsigned int rejected = 0; + auto classify = [&](std::future& result) { + try { + EXPECT_EQ(result.get(), "accepted"); + ++accepted; + } catch (const std::logic_error&) { + ++rejected; + } catch (const std::exception& error) { + ADD_FAILURE() << "Unexpected read failure: " << error.what(); + } + }; + classify(first); + classify(second); + EXPECT_EQ(accepted, 1U); + EXPECT_EQ(rejected, 1U); +} + +TEST_F(StdioTransportTest, PendingWriteOwnsStateAfterTransportDestruction) { + std::istringstream input; + std::ostringstream output; + std::optional> pending; + { + mcp::StdioTransport transport(io_ctx_.get_executor(), input, output); + std::string message = "owned payload"; + pending.emplace(transport.write_message(message)); + message.assign("mutated"); + } + + auto completion = boost::asio::co_spawn(io_ctx_, std::move(*pending), boost::asio::use_future); + io_ctx_.run(); + EXPECT_THROW(completion.get(), std::runtime_error); + EXPECT_TRUE(output.str().empty()); +} + +TEST_F(StdioTransportTest, DestructionWaitsForAnActiveWrite) { + std::istringstream input; + BlockingOutputBuffer output_buffer; + std::ostream output(&output_buffer); + auto transport = std::make_unique(io_ctx_.get_executor(), input, output); + + auto completion = + boost::asio::co_spawn(io_ctx_, transport->write_message("in flight"), boost::asio::use_future); + std::thread runner([this]() { io_ctx_.run(); }); + const bool write_started = output_buffer.wait_until_write_starts(std::chrono::seconds(2)); + if (!write_started) { + output_buffer.release_write(); + runner.join(); + transport.reset(); + FAIL() << "Timed out waiting for the write to start"; + return; + } + + std::atomic destruction_finished{false}; + std::thread destroyer([&]() { + transport.reset(); + destruction_finished.store(true, std::memory_order_release); + }); + std::this_thread::sleep_for(std::chrono::milliseconds(20)); + EXPECT_FALSE(destruction_finished.load(std::memory_order_acquire)); + + output_buffer.release_write(); + destroyer.join(); + runner.join(); + EXPECT_TRUE(destruction_finished.load(std::memory_order_acquire)); + EXPECT_NO_THROW(completion.get()); + EXPECT_EQ(output_buffer.str(), "in flight\n"); +} + +TEST_F(StdioTransportTest, PendingReadOwnsStateAfterTransportDestruction) { + std::istringstream input; + std::ostringstream output; + std::optional> pending; + { + mcp::StdioTransport transport(io_ctx_.get_executor(), input, output); + pending.emplace(transport.read_message()); + } + + auto completion = boost::asio::co_spawn(io_ctx_, std::move(*pending), boost::asio::use_future); + io_ctx_.run(); + EXPECT_THROW(completion.get(), std::runtime_error); +} + TEST_F(StdioTransportTest, CloseIsIdempotent) { std::istringstream input; std::ostringstream output; @@ -143,3 +481,219 @@ TEST_F(StdioTransportTest, PolymorphicThroughBasePointer) { transport->close(); } + +// The protocol channel on this transport is a real process stream, so the tests +// that distinguish a shared stdout from an owned one have to work at the file +// descriptor level. POSIX only; the Windows implementation of +// create_owning_stdout() uses the CRT equivalents and is not exercised here. +#if !defined(_WIN32) + +#include +#include + +namespace { + +// Points the process's stdout and stderr at files the test can read back, and +// puts them back where it found them however the test leaves. +class CapturedProcessStreams { + public: + CapturedProcessStreams(const std::filesystem::path& out, const std::filesystem::path& err) + : saved_stdout_(::dup(STDOUT_FILENO)), saved_stderr_(::dup(STDERR_FILENO)) { + std::cout.flush(); + std::cerr.flush(); + std::fflush(nullptr); + redirect(STDOUT_FILENO, out); + redirect(STDERR_FILENO, err); + } + + ~CapturedProcessStreams() { restore(); } + + CapturedProcessStreams(const CapturedProcessStreams&) = delete; + CapturedProcessStreams& operator=(const CapturedProcessStreams&) = delete; + + void restore() { + std::cout.flush(); + std::cerr.flush(); + std::fflush(nullptr); + if (saved_stdout_ >= 0) { + ::dup2(saved_stdout_, STDOUT_FILENO); + ::close(saved_stdout_); + saved_stdout_ = -1; + } + if (saved_stderr_ >= 0) { + ::dup2(saved_stderr_, STDERR_FILENO); + ::close(saved_stderr_); + saved_stderr_ = -1; + } + } + + private: + static void redirect(int target, const std::filesystem::path& path) { + const int file = ::open(path.c_str(), O_WRONLY | O_CREAT | O_TRUNC, 0600); + ASSERT_GE(file, 0); + ASSERT_GE(::dup2(file, target), 0); + ::close(file); + } + + int saved_stdout_; + int saved_stderr_; +}; + +std::string read_file(const std::filesystem::path& path) { + std::ifstream file(path, std::ios::binary); + return std::string(std::istreambuf_iterator(file), std::istreambuf_iterator()); +} + +// What a peer speaking JSON-RPC over the pipe would have to parse. +std::vector protocol_lines(const std::filesystem::path& path) { + std::vector lines; + std::istringstream contents(read_file(path)); + std::string line; + while (std::getline(contents, line)) { + if (!line.empty()) { + lines.push_back(line); + } + } + return lines; +} + +constexpr const char* kFramedMessage = R"({"jsonrpc":"2.0","id":1,"result":{}})"; +constexpr const char* kStrayPrintf = "libfoo: connected to database"; +constexpr const char* kStrayCout = "cache warmed in 12ms"; + +// Everything an application might innocently put on stdout while the session is +// running. +void emit_stray_application_output() { + std::printf("%s\n", kStrayPrintf); + std::fflush(stdout); + std::cout << kStrayCout << std::endl; +} + +std::filesystem::path scratch_file(const char* name) { + return std::filesystem::temp_directory_path() / name; +} + +} // namespace + +// The footgun itself: the default std::cout output shares the protocol +// channel with the rest of the process. +TEST_F(StdioTransportTest, DefaultOutputSharesTheProtocolChannelWithTheApplication) { + const auto out = scratch_file("mcp_stdio_shared_out.txt"); + const auto err = scratch_file("mcp_stdio_shared_err.txt"); + + { + CapturedProcessStreams capture(out, err); + std::istringstream input; + mcp::StdioTransport transport(io_ctx_.get_executor(), input, std::cout); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await transport.write_message(kFramedMessage); + emit_stray_application_output(); + co_await transport.write_message(kFramedMessage); + }, + boost::asio::detached); + io_ctx_.run(); + std::cout.flush(); + } + + const auto lines = protocol_lines(out); + ASSERT_EQ(lines.size(), 4U) << "expected the two framed messages plus two stray lines"; + EXPECT_EQ(lines[0], kFramedMessage); + EXPECT_EQ(lines[1], kStrayPrintf); + EXPECT_EQ(lines[2], kStrayCout); + EXPECT_EQ(lines[3], kFramedMessage); + + std::filesystem::remove(out); + std::filesystem::remove(err); +} + +// The regression test: an owning transport keeps the channel to itself. +TEST_F(StdioTransportTest, OwnedStdoutKeepsStrayApplicationOutputOffTheProtocolChannel) { + const auto out = scratch_file("mcp_stdio_owned_out.txt"); + const auto err = scratch_file("mcp_stdio_owned_err.txt"); + + { + CapturedProcessStreams capture(out, err); + std::istringstream input; + auto transport = mcp::StdioTransport::create_owning_stdout(io_ctx_.get_executor(), input); + + boost::asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + co_await transport->write_message(kFramedMessage); + emit_stray_application_output(); + co_await transport->write_message(kFramedMessage); + }, + boost::asio::detached); + io_ctx_.run(); + std::cout.flush(); + transport.reset(); + } + + const auto lines = protocol_lines(out); + ASSERT_EQ(lines.size(), 2U) << "protocol channel carried: " << read_file(out); + EXPECT_EQ(lines[0], kFramedMessage); + EXPECT_EQ(lines[1], kFramedMessage); + + // The application's output is not lost, only moved to the diagnostics stream. + const auto diagnostics = read_file(err); + EXPECT_NE(diagnostics.find(kStrayPrintf), std::string::npos); + EXPECT_NE(diagnostics.find(kStrayCout), std::string::npos); + + std::filesystem::remove(out); + std::filesystem::remove(err); +} + +// Destroying the owning transport has to give the process its stdout back, +// otherwise a short-lived session silently swallows everything that follows. +TEST_F(StdioTransportTest, OwnedStdoutIsRestoredWhenTheTransportIsDestroyed) { + const auto out = scratch_file("mcp_stdio_restore_out.txt"); + const auto err = scratch_file("mcp_stdio_restore_err.txt"); + + { + CapturedProcessStreams capture(out, err); + std::istringstream input; + { + auto transport = mcp::StdioTransport::create_owning_stdout(io_ctx_.get_executor(), input); + } + std::cout << "after the session" << std::endl; + } + + EXPECT_NE(read_file(out).find("after the session"), std::string::npos); + + std::filesystem::remove(out); + std::filesystem::remove(err); +} + +// A second owner would duplicate the already-redirected stdout and publish the +// protocol onto stderr, so the attempt has to fail loudly instead. +TEST_F(StdioTransportTest, OnlyOneTransportMayOwnStdoutAtATime) { + const auto out = scratch_file("mcp_stdio_single_out.txt"); + const auto err = scratch_file("mcp_stdio_single_err.txt"); + + // Assertions live outside the capture, otherwise their diagnostics are + // written to the captured file and thrown away with it. + bool second_owner_rejected = false; + std::string rejection; + { + CapturedProcessStreams capture(out, err); + std::istringstream input; + auto first = mcp::StdioTransport::create_owning_stdout(io_ctx_.get_executor(), input); + try { + auto second = mcp::StdioTransport::create_owning_stdout(io_ctx_.get_executor(), input); + } catch (const std::runtime_error& error) { + second_owner_rejected = true; + rejection = error.what(); + } + } + + EXPECT_TRUE(second_owner_rejected) << "a second owner would publish the protocol onto stderr"; + EXPECT_NE(rejection.find("standard output"), std::string::npos) << rejection; + + std::filesystem::remove(out); + std::filesystem::remove(err); +} + +#endif // !defined(_WIN32) diff --git a/test/transport/transport_ws_test.cpp b/test/transport/transport_ws_test.cpp index 1c4782f..910f881 100644 --- a/test/transport/transport_ws_test.cpp +++ b/test/transport/transport_ws_test.cpp @@ -2,19 +2,137 @@ #include +#include "../support/resolve_gate.hpp" +#include "../support/stalling_server.hpp" + +#include +#include +#include +#include #include #include #include #include +#include +#include +#include #include +#include +#include +#include +#include +#include +#include #include +#include #include +#include +#include #include namespace { namespace asio = boost::asio; using tcp = asio::ip::tcp; +using namespace std::chrono_literals; + +struct WriteCompletionGroup { + WriteCompletionGroup(const asio::any_io_executor& executor, std::size_t count) + : signal(executor), remaining(count) { + signal.expires_at(std::chrono::steady_clock::time_point::max()); + } + + void complete(std::exception_ptr operation_error) { + if (operation_error && !error) { + error = operation_error; + } + if (--remaining == 0) { + boost::system::error_code ignored; + signal.cancel(ignored); + } + } + + mcp::Task wait() { + if (remaining != 0) { + boost::system::error_code error_code; + co_await signal.async_wait(asio::redirect_error(asio::use_awaitable, error_code)); + if (error_code && error_code != asio::error::operation_aborted) { + throw boost::system::system_error(error_code); + } + } + if (error) { + std::rethrow_exception(error); + } + } + + asio::steady_timer signal; + std::size_t remaining; + std::exception_ptr error; +}; + +struct ReadCompletion { + explicit ReadCompletion(const asio::any_io_executor& executor) : signal(executor) { + signal.expires_at(std::chrono::steady_clock::time_point::max()); + } + + void complete(std::exception_ptr operation_error, std::string operation_result) { + error = operation_error; + result = std::move(operation_result); + done = true; + boost::system::error_code ignored; + signal.cancel(ignored); + } + + mcp::Task wait() { + if (done) { + co_return; + } + boost::system::error_code error_code; + co_await signal.async_wait(asio::redirect_error(asio::use_awaitable, error_code)); + if (error_code && error_code != asio::error::operation_aborted) { + throw boost::system::system_error(error_code); + } + } + + asio::steady_timer signal; + std::optional result; + std::exception_ptr error; + bool done{false}; +}; + +struct VoidCompletion { + explicit VoidCompletion(const asio::any_io_executor& executor) : signal(executor) { + signal.expires_at(std::chrono::steady_clock::time_point::max()); + } + + void complete(std::exception_ptr operation_error) { + error = operation_error; + done = true; + boost::system::error_code ignored; + signal.cancel(ignored); + } + + mcp::Task wait() { + if (done) { + co_return; + } + boost::system::error_code error_code; + co_await signal.async_wait(asio::redirect_error(asio::use_awaitable, error_code)); + if (error_code && error_code != asio::error::operation_aborted) { + throw boost::system::system_error(error_code); + } + } + + asio::steady_timer signal; + std::exception_ptr error; + bool done{false}; +}; + +mcp::Task wait_for(std::chrono::steady_clock::duration duration) { + auto executor = co_await asio::this_coro::executor; + asio::steady_timer timer(executor, duration); + co_await timer.async_wait(asio::use_awaitable); +} mcp::Task echo_server(tcp::acceptor& acceptor, int echo_count) { auto socket = co_await acceptor.async_accept(asio::use_awaitable); @@ -30,6 +148,17 @@ mcp::Task echo_server(tcp::acceptor& acceptor, int echo_count) { class WebSocketTransportTest : public ::testing::Test { protected: + void run_on_threads(std::size_t count = 4) { + std::vector threads; + threads.reserve(count); + for (std::size_t i = 0; i < count; ++i) { + threads.emplace_back([this]() { io_ctx_.run(); }); + } + for (auto& thread : threads) { + thread.join(); + } + } + asio::io_context io_ctx_; }; @@ -172,3 +301,418 @@ TEST_F(WebSocketTransportTest, ServerTransportAcceptsRawSocket) { EXPECT_EQ(result, "from client"); } + +TEST_F(WebSocketTransportTest, SerializesConcurrentClientWritesAcrossThreads) { + constexpr std::size_t message_count = 32; + tcp::acceptor acceptor(io_ctx_, tcp::endpoint(tcp::v4(), 0)); + auto port = acceptor.local_endpoint().port(); + + std::vector received; + auto server = asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + mcp::WebSocketServerTransport transport(std::move(socket)); + for (std::size_t i = 0; i < message_count; ++i) { + received.push_back(co_await transport.read_message()); + } + transport.close(); + }, + asio::use_future); + + auto client = std::make_shared(io_ctx_.get_executor(), "127.0.0.1", + std::to_string(port)); + std::vector> writes; + writes.reserve(message_count); + for (std::size_t i = 0; i < message_count; ++i) { + writes.push_back(asio::co_spawn(io_ctx_, client->write_message("message-" + std::to_string(i)), + asio::use_future)); + } + + run_on_threads(); + for (auto& write : writes) { + EXPECT_NO_THROW(write.get()); + } + EXPECT_NO_THROW(server.get()); + + std::sort(received.begin(), received.end()); + std::vector expected; + expected.reserve(message_count); + for (std::size_t i = 0; i < message_count; ++i) { + expected.push_back("message-" + std::to_string(i)); + } + std::sort(expected.begin(), expected.end()); + EXPECT_EQ(received, expected); + client->close(); +} + +TEST_F(WebSocketTransportTest, SerializesConcurrentServerWritesDuringHandshake) { + constexpr std::size_t message_count = 32; + tcp::acceptor acceptor(io_ctx_, tcp::endpoint(tcp::v4(), 0)); + auto port = acceptor.local_endpoint().port(); + auto coordinator = asio::make_strand(io_ctx_); + + auto server = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + auto transport = std::make_shared(std::move(socket)); + auto writes = std::make_shared(coordinator, message_count); + for (std::size_t i = 0; i < message_count; ++i) { + asio::co_spawn( + coordinator, transport->write_message("server-" + std::to_string(i)), + asio::bind_executor(coordinator, [writes, transport](std::exception_ptr error) { + writes->complete(error); + })); + } + + EXPECT_EQ(co_await transport->read_message(), "ack"); + co_await writes->wait(); + transport->close(); + }, + asio::use_future); + + std::vector received; + auto client = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + mcp::WebSocketClientTransport transport(io_ctx_.get_executor(), "127.0.0.1", + std::to_string(port)); + for (std::size_t i = 0; i < message_count; ++i) { + received.push_back(co_await transport.read_message()); + } + co_await transport.write_message("ack"); + transport.close(); + }, + asio::use_future); + + run_on_threads(); + EXPECT_NO_THROW(client.get()); + EXPECT_NO_THROW(server.get()); + + std::sort(received.begin(), received.end()); + std::vector expected; + expected.reserve(message_count); + for (std::size_t i = 0; i < message_count; ++i) { + expected.push_back("server-" + std::to_string(i)); + } + std::sort(expected.begin(), expected.end()); + EXPECT_EQ(received, expected); +} + +TEST_F(WebSocketTransportTest, RejectsSecondOutstandingClientRead) { + tcp::acceptor acceptor(io_ctx_, tcp::endpoint(tcp::v4(), 0)); + auto port = acceptor.local_endpoint().port(); + auto coordinator = asio::make_strand(io_ctx_); + + auto server = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + mcp::WebSocketServerTransport transport(std::move(socket)); + EXPECT_EQ(co_await transport.read_message(), "ready"); + co_await wait_for(40ms); + co_await transport.write_message("response"); + transport.close(); + }, + asio::use_future); + + bool rejected = false; + auto client = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + mcp::WebSocketClientTransport transport(io_ctx_.get_executor(), "127.0.0.1", + std::to_string(port)); + co_await transport.write_message("ready"); + + auto first_read = std::make_shared(coordinator); + asio::co_spawn(coordinator, transport.read_message(), + asio::bind_executor( + coordinator, [first_read](std::exception_ptr error, std::string result) { + first_read->complete(error, std::move(result)); + })); + co_await wait_for(5ms); + try { + (void)co_await transport.read_message(); + } catch (const std::logic_error&) { + rejected = true; + } + + co_await first_read->wait(); + if (first_read->error) { + std::rethrow_exception(first_read->error); + } + EXPECT_EQ(first_read->result, "response"); + transport.close(); + }, + asio::use_future); + + run_on_threads(); + EXPECT_NO_THROW(client.get()); + EXPECT_NO_THROW(server.get()); + EXPECT_TRUE(rejected); +} + +TEST_F(WebSocketTransportTest, PendingClientReadOwnsStateAfterTransportDestruction) { + tcp::acceptor acceptor(io_ctx_, tcp::endpoint(tcp::v4(), 0)); + auto port = acceptor.local_endpoint().port(); + auto coordinator = asio::make_strand(io_ctx_); + + auto server = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + mcp::WebSocketServerTransport transport(std::move(socket)); + EXPECT_EQ(co_await transport.read_message(), "ready"); + co_await wait_for(50ms); + transport.close(); + }, + asio::use_future); + + bool read_failed = false; + auto client = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + auto transport = std::make_unique( + io_ctx_.get_executor(), "127.0.0.1", std::to_string(port)); + co_await transport->write_message("ready"); + + auto pending_read = std::make_shared(coordinator); + asio::co_spawn(coordinator, transport->read_message(), + asio::bind_executor(coordinator, [pending_read](std::exception_ptr error, + std::string result) { + pending_read->complete(error, std::move(result)); + })); + co_await wait_for(5ms); + transport.reset(); + co_await pending_read->wait(); + read_failed = pending_read->error != nullptr; + }, + asio::use_future); + + run_on_threads(); + EXPECT_NO_THROW(client.get()); + EXPECT_NO_THROW(server.get()); + EXPECT_TRUE(read_failed); +} + +TEST_F(WebSocketTransportTest, PendingServerReadOwnsStateAfterTransportDestruction) { + tcp::acceptor acceptor(io_ctx_, tcp::endpoint(tcp::v4(), 0)); + auto port = acceptor.local_endpoint().port(); + auto coordinator = asio::make_strand(io_ctx_); + + bool read_failed = false; + auto server = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + auto transport = std::make_unique(std::move(socket)); + EXPECT_EQ(co_await transport->read_message(), "ready"); + + auto pending_read = std::make_shared(coordinator); + asio::co_spawn(coordinator, transport->read_message(), + asio::bind_executor(coordinator, [pending_read](std::exception_ptr error, + std::string result) { + pending_read->complete(error, std::move(result)); + })); + co_await wait_for(5ms); + transport.reset(); + co_await pending_read->wait(); + read_failed = pending_read->error != nullptr; + }, + asio::use_future); + + auto client = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + mcp::WebSocketClientTransport transport(io_ctx_.get_executor(), "127.0.0.1", + std::to_string(port)); + co_await transport.write_message("ready"); + co_await wait_for(50ms); + transport.close(); + }, + asio::use_future); + + run_on_threads(); + EXPECT_NO_THROW(client.get()); + EXPECT_NO_THROW(server.get()); + EXPECT_TRUE(read_failed); +} + +TEST_F(WebSocketTransportTest, CanceledQueuedWriteDoesNotBlockSubsequentWrites) { + constexpr std::size_t payload_size = 16 * 1024 * 1024; + tcp::acceptor acceptor(io_ctx_, tcp::endpoint(tcp::v4(), 0)); + auto port = acceptor.local_endpoint().port(); + auto coordinator = asio::make_strand(io_ctx_); + + auto server = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + auto socket = co_await acceptor.async_accept(asio::use_awaitable); + mcp::WebSocketServerTransport transport(std::move(socket)); + co_await transport.write_message("ready"); + co_await wait_for(100ms); + + auto large_message = co_await transport.read_message(); + EXPECT_EQ(large_message.size(), payload_size); + EXPECT_EQ(co_await transport.read_message(), "after-cancel"); + co_await transport.write_message("done"); + transport.close(); + }, + asio::use_future); + + bool canceled_write_failed = false; + auto client = asio::co_spawn( + coordinator, + [&]() -> mcp::Task { + auto transport = std::make_shared( + io_ctx_.get_executor(), "127.0.0.1", std::to_string(port)); + EXPECT_EQ(co_await transport->read_message(), "ready"); + + auto watchdog = std::make_shared(coordinator, 30s); + watchdog->async_wait( + asio::bind_executor(coordinator, [transport](const boost::system::error_code& error) { + if (!error) { + transport->close(); + } + })); + + auto first_write = std::make_shared(coordinator); + asio::co_spawn(coordinator, transport->write_message(std::string(payload_size, 'x')), + asio::bind_executor(coordinator, [first_write](std::exception_ptr error) { + first_write->complete(error); + })); + co_await wait_for(10ms); + + asio::cancellation_signal cancel_queued_write; + auto canceled_write = std::make_shared(coordinator); + asio::co_spawn( + coordinator, transport->write_message("cancel-me"), + asio::bind_cancellation_slot( + cancel_queued_write.slot(), + asio::bind_executor(coordinator, [canceled_write](std::exception_ptr error) { + canceled_write->complete(error); + }))); + co_await wait_for(10ms); + cancel_queued_write.emit(asio::cancellation_type::all); + co_await canceled_write->wait(); + canceled_write_failed = canceled_write->error != nullptr; + + co_await transport->write_message("after-cancel"); + co_await first_write->wait(); + if (first_write->error) { + std::rethrow_exception(first_write->error); + } + EXPECT_EQ(co_await transport->read_message(), "done"); + + boost::system::error_code ignored; + watchdog->cancel(ignored); + transport->close(); + }, + asio::use_future); + + run_on_threads(); + EXPECT_NO_THROW(client.get()); + EXPECT_NO_THROW(server.get()); + EXPECT_TRUE(canceled_write_failed); +} + +// A peer that accepts the connection and then never answers the upgrade request must not hold a +// pending call forever: the connect timeout ends it. +TEST_F(WebSocketTransportTest, StalledHandshakeFailsAfterTheConnectTimeout) { + StallingServer stalling(io_ctx_); + stalling.accept_and_stall(); + + mcp::WebSocketClientTransport transport(io_ctx_.get_executor(), "127.0.0.1", + std::to_string(stalling.port()), "/", + std::chrono::milliseconds(200)); + + auto result = asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + try { + co_await transport.write_message(R"({"jsonrpc":"2.0","id":1,"method":"ping"})"); + } catch (const boost::system::system_error& error) { + co_return error.code() == boost::beast::error::timeout; + } + co_return false; + }, + asio::use_future); + + std::thread runner([this]() { io_ctx_.run(); }); + const bool finished = result.wait_for(std::chrono::seconds(10)) == std::future_status::ready; + io_ctx_.stop(); + runner.join(); + + ASSERT_TRUE(finished) << "the stalled handshake was never cut off"; + EXPECT_TRUE(result.get()) << "the pending call must fail with a timeout"; + EXPECT_EQ(stalling.accepted(), 1); +} + +#ifdef __linux__ + +// close() called from a thread that does not run the io_context, while the connect is suspended +// in its resolve. Asio's resolver thread checks its cancel token once, before it calls +// getaddrinfo(), so a close() that arrives while that call is in progress finds a resolve it can +// no longer cancel and a socket that is not open yet. The gate holds the lookup inside +// getaddrinfo() until close() has done its cancelling, so the connect resumes with usable addresses +// on a transport that is already closed. It must stop there. Without the closed check in +// ensure_connected() it connects to the stalling server and waits on a handshake that never ends. +TEST_F(WebSocketTransportTest, CloseWhileTheResolverIsPastItsCancelCheckOpensNoConnection) { + StallingServer stalling(io_ctx_); + stalling.accept_and_stall(); + + auto transport = std::make_shared( + io_ctx_.get_executor(), "127.0.0.1", std::to_string(stalling.port())); + + resolve_gate().arm(stalling.port()); + + std::exception_ptr failure; + std::promise write_done; + auto write_finished = write_done.get_future(); + asio::co_spawn( + io_ctx_, + [&]() -> mcp::Task { + try { + co_await transport->write_message(R"({"jsonrpc":"2.0","id":1,"method":"ping"})"); + } catch (...) { + failure = std::current_exception(); + } + write_done.set_value(); + }, + asio::detached); + + std::thread runner([this]() { io_ctx_.run(); }); + + const auto limit = std::chrono::seconds(10); + const bool lookup_held = resolve_gate().wait_until_entered(limit); + + // close() posts its cancelling to the transport's strand, which is one hop through the + // io_context. Two further hops through the same queue cannot complete before it has. + std::promise cancel_done; + auto cancel_finished = cancel_done.get_future(); + bool cancel_ran = false; + if (lookup_held) { + transport->close(); + asio::post(io_ctx_, [&]() { asio::post(io_ctx_, [&]() { cancel_done.set_value(); }); }); + cancel_ran = cancel_finished.wait_for(limit) == std::future_status::ready; + } + + resolve_gate().release(); + const bool finished = write_finished.wait_for(limit) == std::future_status::ready; + + // Everything below reads state the runner thread wrote, so it stops first. + io_ctx_.stop(); + runner.join(); + + ASSERT_TRUE(lookup_held) << "the connect never reached getaddrinfo()"; + ASSERT_TRUE(cancel_ran) << "the cancelling posted by close() never ran on the io_context"; + EXPECT_TRUE(finished) << "close() was lost: the connect resumed from the lookup after the " + "transport closed and is still running"; + EXPECT_EQ(stalling.accepted(), 0) << "the transport connected after it had been closed"; + if (finished) { + EXPECT_NE(failure, nullptr) << "a write cut short by close() must report an error"; + } +} + +#endif // __linux__ diff --git a/test/tsan.supp b/test/tsan.supp new file mode 100644 index 0000000..3e6214f --- /dev/null +++ b/test/tsan.supp @@ -0,0 +1,17 @@ +# ThreadSanitizer suppressions for the test suite. Each entry names a report that is not a defect in +# this SDK; keep the list this short and add nothing without the reasoning beside it. +# +# A suppression discards a report after ThreadSanitizer has built it, so it hides the text and not +# the cost: every thread still stops while the stacks are symbolized. + +# Asio frees a shared executor after a release decrement of its reference count followed by +# atomic_thread_fence(acquire), the standard pattern. ThreadSanitizer still reports the delete as +# racing with the decrement, under both GCC 13 and Clang 18, so it does not credit the standalone +# fence. Seen wherever co_spawn() runs on a strand. +race:boost::asio::execution::detail::shared_target_executor + +# Asio's signal handler reads the write end of its self-pipe, which the signal_set_service +# constructor wrote on another thread. The tests raise the signal after a sleep that follows the +# construction, an ordering ThreadSanitizer cannot see; the service is never used before it is +# built. +race:boost_asio_signal_handler